Diff Coverage

Diff: origin/r1.0.0...HEAD, staged and unstaged changes

Source File Diff Coverage (%) Missing Lines
hyper_parallel/core/context_parallel/async_dsa_context_parallel.py 100%  
hyper_parallel/core/context_parallel/context_parallel.py 47.4% 49,863-865,936-937,948-949,951-954,962-963,967-968,974-975,980-981
hyper_parallel/core/context_parallel/dsa_context_parallel.py 72.2% 286,294-295,297-298
hyper_parallel/core/shard/ops/dsa_cp_fold.py 64.6% 133-138,156,169,176-178,207-209,220,250-257,259-265,267-273,275-277,283,286,307-310,316,345,351-352,400,407-408,449,484-487,490,492-496,522-523,525-532,535-537,539-545,594-601,604,606-608,610,630,633,648,652,680,690,693-694,703,761
hyper_parallel/core/shard/ops/parallel_lightning_indexer.py 50.0% 314,316,340,342
hyper_parallel/core/shard/ops/parallel_npu_dense_lightning_indexer_grad_kl_loss.py 33.3% 410-412,417,477,480-481,484
hyper_parallel/core/shard/ops/parallel_npu_dense_lightning_indexer_softmax_lse.py 40.0% 387,390,418,423-424,427
hyper_parallel/core/shard/ops/parallel_npu_flash_attention_score.py 50.0% 84
hyper_parallel/core/shard/ops/parallel_npu_sparse_flash_attention.py 50.0% 413,415,442,447
hyper_parallel/core/shard/ops/parallel_npu_sparse_lightning_indexer_grad_kl_loss.py 33.3% 405-407,412,474,481-482,486
hyper_parallel/platform/mindspore/platform.py 33.3% 385-388,392-395,403-404,408,1280-1282
hyper_parallel/core/context_parallel/context_parallel.py
45
46
47
48
49
50
51
52
53
    Boundaries are free to return one, so unpack for that case.
    """
    items = list(items)
    if isinstance(original, tuple) and hasattr(original, "_fields"):
        return type(original)(*items)
    return type(original)(items)


def _same_mesh_rank_list(lhs: DeviceMesh, rhs: DeviceMesh) -> bool:
859
860
861
862
863
864
865
866
867
868
869
        same problem and cannot be detected here.)
        """
        if platform.platform_type != PlatformType.MINDSPORE:
            return
        import mindspore  # pylint: disable=import-outside-toplevel
        if mindspore.get_context("mode") == mindspore.GRAPH_MODE:
            raise NotImplementedError(
                "Head-tail load balance for Colossal CP is implemented for PyNative only: it "
                "replaces the module's construct on the instance, which graph mode ignores. "
                "Run in PyNative or leave load_balance off."
            )
932
933
934
935
936
937
938
939
940
941
        # K/V are Replicate; wrap once and reuse for both FA calls
        k_full_dt = DTensor.from_local(k_full, co_submesh, (Replicate(),))
        v_full_dt = DTensor.from_local(v_full, co_submesh, (Replicate(),))

        def _localize(value):
            return value.to_local() if isinstance(value, DTensor) else value

        def _fa(q_half, split_id):
            new_args[q_idx] = DTensor.from_local(q_half, co_submesh, (Shard(seq_dim),))
            new_args[k_idx] = k_full_dt
944
945
946
947
948
949
950
951
952
953
954
955
956
957
958
            # MindFormers' TND softmax converter (``current_lb_split``). Leaking it past an
            # exception in the boundary would leave every later call reading this sub-call's
            # chunk id -- which shows up as a wrong loss, not as an error.
            _set_lb_override(split_id=split_id, split_num=2 * ws)
            try:
                out = original_forward(*new_args, **kwargs)
            finally:
                _clear_lb_override()
            if isinstance(out, (tuple, list)):
                return _same_sequence_type(out, (_localize(item) for item in out))
            return _localize(out)

        fa1_out = _fa(q_keep, split_id=2 * local_idx)
        fa2_out = _fa(q_peer, split_id=2 * target_idx + 1)
        # A boundary may return several per-query tensors (the DSA dense teacher hands back
958
959
960
961
962
963
964
965
966
967
968
969
970
971
972
        # A boundary may return several per-query tensors (the DSA dense teacher hands back
        # attention_out plus its softmax statistics); every one of them belongs to the peer's
        # query half and is stitched on the same sequence axis -- MindFormers canonicalises
        # the statistics to the output's axes before they cross this boundary.
        if isinstance(fa1_out, (tuple, list)):
            for index, item in enumerate(fa1_out):
                # Every element must be per-query, or stitching it on the sequence axis is
                # wrong -- silently so: a None would crash deep inside p2p_exchange and a
                # reduced scalar (a loss, say) would be concatenated into nonsense.
                if item is None or not hasattr(item, "shape") or len(item.shape) <= seq_dim:
                    raise NotImplementedError(
                        f"Head-tail load balance needs every output of the attention boundary "
                        f"to be a per-query tensor stitched on dim {seq_dim}, but output "
                        f"{index} is {type(item).__name__}. An output that is not per-query "
                        f"would have to be combined, not concatenated."
970
971
972
973
974
975
976
977
978
979
980
981
982
                        f"to be a per-query tensor stitched on dim {seq_dim}, but output "
                        f"{index} is {type(item).__name__}. An output that is not per-query "
                        f"would have to be combined, not concatenated."
                    )
            fa2_our = [platform.p2p_exchange(item, peer_rank) for item in fa2_out]
            out = _same_sequence_type(fa1_out, (
                platform.cat([first, second], dim=seq_dim)
                for first, second in zip(fa1_out, fa2_our)
            ))
        else:
            fa2_our = platform.p2p_exchange(fa2_out, peer_rank)
            out = platform.cat([fa1_out, fa2_our], dim=seq_dim)
        return _finalize_colossal_output(out, output_layout, co_submesh, seq_dim, self.use_local_output)
hyper_parallel/core/context_parallel/dsa_context_parallel.py
282
283
284
285
286
287
288
289
290
        # model or test in the same process inheriting a stale True -- still works, because
        # then nobody has asked for folding and this is a no-op.
        previous = dsa_cp_fold_requester()
        if previous:
            raise ValueError(
                f"{style_name}(load_balance=False) would turn off DSA CP head-tail folding that "
                f"{previous} already turned on for this process. The DSA styles must agree: pass "
                f"load_balance to every one of them (it defaults to False), or to none. If you "
                f"really want to reset the flag between models, call set_dsa_cp_fold(False) "
290
291
292
293
294
295
296
297
298
299
300
301
302
                f"really want to reset the flag between models, call set_dsa_cp_fold(False) "
                f"explicitly at model teardown.")
        set_dsa_cp_fold(False)
        return False
    if layout not in ("BSND", "TND"):
        raise ValueError(
            f"{style_name}(load_balance=True) supports BSND and TND layouts, got {layout!r}.")
    set_dsa_cp_fold(True, requester=style_name)
    return True


def _query_stats_seq_dim(layout: str) -> int:
    """Return the sequence dimension used by query-side softmax stats."""
hyper_parallel/core/shard/ops/dsa_cp_fold.py
129
130
131
132
133
134
135
136
137
138
139
140
141
142
    Entry ``s`` names the prefix chunk that carries folded slot ``s``'s gradient; slots
    the prefix does not cover point at ``m`` (one past the prefix), which the caller
    fills with a zero chunk, so a single ``index_select`` scatters and zero-fills.
    """
    prefix = build_prefix_order(seq_shard_id, seq_shards, block_id)
    m = len(prefix)
    inverse = [m] * (2 * seq_shards)
    for i, slot in enumerate(prefix):
        inverse[slot] = i
    return inverse


# ---------------------------------------------------------------------------
# Tensor helpers
152
153
154
155
156
157
158
159
160
    if idx is None:
        idx = platform.from_numpy(np.asarray(values, dtype=np.int32))
        _INDEX_CACHE[key] = idx
    if type(ref).__module__.startswith("torch"):
        idx = idx.to(ref.device)
    return idx


def _chunked_shape(shape: tuple, seq_dim: int, chunks: int) -> tuple:
165
166
167
168
169
170
171
172
173
def _unfold_prefix_select(x, seq_shard_id: int, seq_shards: int, block_id: int, seq_dim: int):
    """One folded block's causal key prefix, gathered out of the folded-order full-length key."""
    twon = 2 * seq_shards
    if x.shape[seq_dim] % twon != 0:
        raise ValueError(
            f"DSA CP fold needs the key sequence ({x.shape[seq_dim]}) to be a multiple of 2 * cp ({twon})."
        )
    sf = x.shape[seq_dim] // twon
    prefix = build_prefix_order(seq_shard_id, seq_shards, block_id)
172
173
174
175
176
177
178
179
180
181
182
    sf = x.shape[seq_dim] // twon
    prefix = build_prefix_order(seq_shard_id, seq_shards, block_id)
    y = x.reshape(_chunked_shape(x.shape, seq_dim, twon))
    y = y.index_select(seq_dim, _index_tensor(prefix, x))
    out_shape = list(x.shape)
    out_shape[seq_dim] = len(prefix) * sf
    return y.reshape(tuple(out_shape))


class _UnfoldPrefix(platform.Function):
    """``index_select`` forward whose backward is another ``index_select``, not ``index_add``.
203
204
205
206
207
208
209
210
211
212
213
        ``aclnnIndexAdd`` has no bf16 kernel -- it would cast the whole key gradient
        bf16->fp32->bf16 on every call. The prefix slots do not repeat, so its inverse
        is exactly ``fold_prefix_grad`` (pad one zero chunk + one gather, stays bf16).
        """
        seq_shard_id, seq_shards, block_id, full_len, seq_dim = ctx.fold_args
        grad = fold_prefix_grad(grad_output, seq_shard_id, seq_shards, block_id, full_len, seq_dim)
        return grad, None, None, None, None


def unfold_prefix(x, seq_shard_id: int, seq_shards: int, block_id: int, seq_dim: int = 1):
    """Natural-order causal key prefix of one folded block, from a folded full-length key.
216
217
218
219
220
221
222
223
224
    Differentiable: the gradient goes back onto the folded layout through
    ``fold_prefix_grad`` (see ``_UnfoldPrefix``).
    """
    if x is None:
        return None
    return _UnfoldPrefix.apply(x, seq_shard_id, seq_shards, block_id, seq_dim)


class _UnfoldPrefixPair(platform.Function):
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
    @staticmethod
    def backward(ctx, grad0, grad1):  # pylint: disable=arguments-differ
        """Merged reverse of both prefixes: the shorter one is a head of the longer one,
        so add it into that head and scatter back with a single gather."""
        seq_shard_id, seq_shards, full_len, seq_dim, sf, m = ctx.fold_args
        long_id = 0 if m[0] >= m[1] else 1
        grads = (grad0, grad1)
        g_long, g_short = grads[long_id], grads[1 - long_id]
        if g_long is None and g_short is None:
            return None, None, None, None
        if g_long is None:
            return fold_prefix_grad(g_short, seq_shard_id, seq_shards, 1 - long_id, full_len, seq_dim), \
                None, None, None
        parts = []
        short_len = m[1 - long_id] * sf
        long_len = m[long_id] * sf
        if g_short is not None:
            parts.append(g_long.narrow(seq_dim, 0, short_len) + g_short)
            if long_len > short_len:
                parts.append(g_long.narrow(seq_dim, short_len, long_len - short_len))
        else:
            parts.append(g_long)
        zero_shape = list(g_long.shape)
        zero_shape[seq_dim] = sf
        parts.append(platform.zeros(tuple(zero_shape), dtype=g_long.dtype, device=g_long.device))
        g = platform.cat(parts, dim=seq_dim)
        g = g.reshape(_chunked_shape(g.shape, seq_dim, m[long_id] + 1))
        g = g.index_select(seq_dim, _index_tensor(
            build_prefix_scatter_order(seq_shard_id, seq_shards, long_id), g))
        out_shape = list(g_long.shape)
        out_shape[seq_dim] = full_len
        return g.reshape(tuple(out_shape)), None, None, None


def unfold_prefix_pair(x, seq_shard_id: int, seq_shards: int, seq_dim: int = 1):
    """``(unfold_prefix(x, .., 0), unfold_prefix(x, .., 1))`` with one merged backward."""
279
280
281
282
283
284
285
286
287
288
289

def unfold_prefix_pair(x, seq_shard_id: int, seq_shards: int, seq_dim: int = 1):
    """``(unfold_prefix(x, .., 0), unfold_prefix(x, .., 1))`` with one merged backward."""
    if x is None:
        return None, None
    twon = 2 * seq_shards
    if x.shape[seq_dim] % twon != 0:
        raise ValueError(
            f"DSA CP fold needs the key sequence ({x.shape[seq_dim]}) to be a multiple of 2 * cp ({twon})."
        )
    return _UnfoldPrefixPair.apply(x, seq_shard_id, seq_shards, seq_dim)
303
304
305
306
307
308
309
310
311
312
313
314
    g = grad.reshape(_chunked_shape(grad.shape, seq_dim, m))
    zero_shape = list(g.shape)
    zero_shape[seq_dim] = 1
    g = platform.cat([g, platform.zeros(tuple(zero_shape), dtype=g.dtype, device=g.device)], dim=seq_dim)
    g = g.index_select(seq_dim, _index_tensor(build_prefix_scatter_order(seq_shard_id, seq_shards, block_id), g))
    out_shape = list(grad.shape)
    out_shape[seq_dim] = full_len
    return g.reshape(tuple(out_shape))


def split_half(x, dim: int) -> Tuple[Optional[object], Optional[object]]:
    """Split a query-side tensor into its two folded blocks (views)."""
312
313
314
315
316
317
318
319

def split_half(x, dim: int) -> Tuple[Optional[object], Optional[object]]:
    """Split a query-side tensor into its two folded blocks (views)."""
    if x is None:
        return None, None
    half = x.shape[dim] // 2
    return x.narrow(dim, 0, half), x.narrow(dim, half, x.shape[dim] - half)

341
342
343
344
345
346
347
348
349
    ``C_last < T``, and the last block then fails inside the kernel with a message about
    sequence lengths that says nothing about folding or padding. Fail here instead, once.
    """
    if _KLEN_CHECKED[0]:
        return
    _KLEN_CHECKED[0] = True
    try:
        last = int(actual_seq_klen[-1])
    except Exception:  # pylint: disable=broad-except
347
348
349
350
351
352
353
354
355
356
    try:
        last = int(actual_seq_klen[-1])
    except Exception:  # pylint: disable=broad-except
        return          # cannot read it (traced/placeholder tensor) -- leave it to the kernel
    if last != int(full_len):
        raise ValueError(
            f"DSA CP head-tail fold needs the key cumulative lengths to cover the whole key "
            f"tensor, but actual_seq_klen[-1]={last} while the key length is {int(full_len)}. "
            f"A padded tail (eod_pad_length) is not supported by the folded per-block "
            f"cumulative lengths: the kernel checks that the accumulated key length equals T. "
396
397
398
399
400
401
402
403
404
    Returns:
        tuple: ``(Q_b, K_b)`` as int32 tensors of the same length as the inputs.
    """
    if actual_seq_qlen is None or actual_seq_klen is None:
        return actual_seq_qlen, actual_seq_klen
    _check_klen_covers_full_len(actual_seq_klen, full_len)
    sf = int(full_len) // (2 * int(seq_shards))
    # The block's causal prefix spans this many Sf chunks, so its own chunk ends the prefix.
    end = balanced_prefix_chunks(seq_shard_id, seq_shards, block_id) * sf
403
404
405
406
407
408
409
410
411
412
    # The block's causal prefix spans this many Sf chunks, so its own chunk ends the prefix.
    end = balanced_prefix_chunks(seq_shard_id, seq_shards, block_id) * sf
    start = end - sf
    q_block = platform.tensor_type_cast((actual_seq_qlen - start).clamp(0, sf), 'int32')
    k_block = platform.tensor_type_cast(actual_seq_klen.clamp(0, end), 'int32')
    return q_block, k_block


# ---------------------------------------------------------------------------
# Per-kernel twice-call wrappers (BSND). ``func`` is the local kernel callable the
445
446
447
448
449
450
451
452
453
    The invariant: the local query holds exactly the two folded chunks this rank owns
    (``2 * sf``) while the key holds all ``2 * N`` of them, hence ``local_q_len * N == key_len``.
    """
    if local_q is None or key is None:
        return
    q_len = int(local_q.shape[seq_dim])
    k_len = int(key.shape[seq_dim])
    if q_len * int(seq_shards) != k_len:
        raise ValueError(
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
    key = args[1]
    rest = args[3:]
    k0, k1 = unfold_prefix_pair(key, seq_shard_id, seq_shards, seq_dim)

    def _kwargs(block_id):
        if not tnd:
            return kwargs
        q_len, k_len = tnd_block_seq_lens(
            kwargs.get("actual_seq_lengths_query"), kwargs.get("actual_seq_lengths_key"),
            key.shape[0], seq_shard_id, seq_shards, block_id)
        return {**kwargs, "actual_seq_lengths_query": q_len, "actual_seq_lengths_key": k_len}

    out0 = func(q0, k0, w0, *rest, **_kwargs(0))
    out1 = func(q1, k1, w1, *rest, **_kwargs(1))
    if not isinstance(out0, (tuple, list)):
        return _cat_pair(out0, out1, seq_dim)
    return _same_sequence_type(out0, (_cat_pair(a, b, seq_dim) for a, b in zip(out0, out1)))


def fold_sparse_flash_attention(func: Callable, seq_shard_id: int, seq_shards: int, *args,
                                fold_layout: str = "BSND", **kwargs):
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
    qr0, qr1 = split_half(kwargs.get("query_rope"), seq_dim)
    key_rope = kwargs.get("key_rope")

    keys = unfold_prefix_pair(key, seq_shard_id, seq_shards, seq_dim)
    values = unfold_prefix_pair(value, seq_shard_id, seq_shards, seq_dim)
    key_ropes = unfold_prefix_pair(key_rope, seq_shard_id, seq_shards, seq_dim)

    def _call(block_id, q, topk, q_rope):
        kw = dict(kwargs)
        if "query_rope" in kw:
            kw["query_rope"] = q_rope
        if "key_rope" in kw:
            kw["key_rope"] = key_ropes[block_id]
        if tnd:
            q_len, k_len = tnd_block_seq_lens(
                kw.get("actual_seq_lengths_query"), kw.get("actual_seq_lengths_kv"),
                key.shape[0], seq_shard_id, seq_shards, block_id)
            kw["actual_seq_lengths_query"] = q_len
            kw["actual_seq_lengths_kv"] = k_len
        return func(q, keys[block_id], values[block_id], topk, *rest, **kw)

    out0 = _call(0, q0, t0, qr0)
    out1 = _call(1, q1, t1, qr1)
    if not isinstance(out0, (tuple, list)):
        return _cat_pair(out0, out1, seq_dim)
    stitched = [_cat_pair(out0[0], out1[0], seq_dim)]
    stitched.extend(_cat_pair(a, b, stats_dim) for a, b in zip(out0[1:], out1[1:]))
    return _same_sequence_type(out0, stitched)


# ``args`` positions of the MindSpore positional form of the sparse indexer KL loss:
#   0 query, 1 key, 2 query_index, 3 key_index, 4 weights, 5 sparse_indices,
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
    full_len = args[3].shape[seq_dim]
    halves = {i: split_half(args[i], d) for i, d in q_side_dims.items()}
    prefixes = {i: unfold_prefix_pair(args[i], seq_shard_id, seq_shards, seq_dim) for i in key_side}

    def _block(block_id):
        call = list(args)
        for i in q_side_dims:
            call[i] = halves[i][block_id]
        for i in key_side:
            call[i] = prefixes[i][block_id]
        if tnd:
            call[_KL_ACTUAL_SEQ_QLEN_IDX], call[_KL_ACTUAL_SEQ_KLEN_IDX] = tnd_block_seq_lens(
                args[_KL_ACTUAL_SEQ_QLEN_IDX], args[_KL_ACTUAL_SEQ_KLEN_IDX],
                full_len, seq_shard_id, seq_shards, block_id)
        return func(*call)

    d_qi0, d_ki0, d_w0, loss0 = _block(0)
    d_qi1, d_ki1, d_w1, loss1 = _block(1)
    d_key_index = (fold_prefix_grad(d_ki0, seq_shard_id, seq_shards, 0, full_len, seq_dim)
                   + fold_prefix_grad(d_ki1, seq_shard_id, seq_shards, 1, full_len, seq_dim))
    return (_cat_pair(d_qi0, d_qi1, seq_dim), d_key_index,
            _cat_pair(d_w0, d_w1, seq_dim), loss0 + loss1)


626
627
628
629
630
631
632
633
634
635
636
637
    whose causal prefix is a *contiguous* run from the start -- so this is a narrow, not
    the ``index_select`` the folded-key path needs.
    """
    if x is None:
        return None
    twon = 2 * seq_shards
    if x.shape[seq_dim] % twon != 0:
        raise ValueError(
            f"DSA CP fold needs the key sequence ({x.shape[seq_dim]}) to be a multiple of 2 * cp ({twon})."
        )
    sf = x.shape[seq_dim] // twon
    return x.narrow(seq_dim, 0, balanced_prefix_chunks(seq_shard_id, seq_shards, block_id) * sf)
644
645
646
647
648
649
650
651
652
653
654
655
656
    has to be scattered: the positions the prefix did not cover are exactly the tail.
    """
    pad_len = full_len - grad.shape[seq_dim]
    if pad_len <= 0:
        return grad
    pad_shape = list(grad.shape)
    pad_shape[seq_dim] = pad_len
    zero = platform.zeros(tuple(pad_shape), dtype=grad.dtype, device=grad.device)
    return platform.cat([grad, zero], dim=seq_dim)


# ``args`` positions of the MindSpore positional form of the dense indexer forward:
#   0 query_index, 1 key_index, 2 weights, 3 actual_seq_qlen, 4 actual_seq_klen,
676
677
678
679
680
681
682
683
684
    tnd = fold_layout == "TND"
    seq_dim = 0 if tnd else 1
    stats_dim = 1 if tnd else 2
    if tnd and len(args) <= _DENSE_LSE_ACTUAL_SEQ_KLEN_IDX:
        raise NotImplementedError(
            "DSA CP head-tail fold under TND needs the MindSpore positional signature of the "
            "dense indexer forward, which carries actual_seq_qlen/klen.")
    check_fold_shapes(args[0], args[1], seq_shards, seq_dim)
    q0, q1 = split_half(args[0], seq_dim)
686
687
688
689
690
691
692
693
694
695
696
697
698

    def _call(block_id, query_index, weights):
        rest = list(args[3:])
        if tnd:
            q_len, k_len = tnd_block_seq_lens(
                args[_DENSE_LSE_ACTUAL_SEQ_QLEN_IDX], args[_DENSE_LSE_ACTUAL_SEQ_KLEN_IDX],
                args[1].shape[seq_dim], seq_shard_id, seq_shards, block_id)
            rest[_DENSE_LSE_ACTUAL_SEQ_QLEN_IDX - 3] = q_len
            rest[_DENSE_LSE_ACTUAL_SEQ_KLEN_IDX - 3] = k_len
        key_prefix = natural_prefix(args[1], seq_shard_id, seq_shards, block_id, seq_dim)
        return func(query_index, key_prefix, weights, *rest, **kwargs)

    out0 = _call(0, q0, w0)
699
700
701
702
703
704
705
706
707
    out1 = _call(1, q1, w1)
    if not isinstance(out0, (tuple, list)):
        # Mirrors the guard the sibling wrappers already have: a single-tensor return is
        # stitched directly rather than zipped element-wise.
        return _cat_pair(out0, out1, stats_dim)
    return _same_sequence_type(out0, (_cat_pair(a, b, stats_dim) for a, b in zip(out0, out1)))


# ``args`` positions of the MindSpore positional form of the dense indexer KL loss:
757
758
759
760
761
762
    d_qi0, d_ki0, d_w0, loss0 = _block(0)
    d_qi1, d_ki1, d_w1, loss1 = _block(1)
    d_key_index = (natural_prefix_grad(d_ki0, full_len, seq_dim)
                   + natural_prefix_grad(d_ki1, full_len, seq_dim))
    return (_cat_pair(d_qi0, d_qi1, seq_dim), d_key_index,
            _cat_pair(d_w0, d_w1, seq_dim), loss0 + loss1)
hyper_parallel/core/shard/ops/parallel_lightning_indexer.py
310
311
312
313
314
315
316
317
318
319
            split_id = q_layout.get_split_id(1)
            seq_shards = q_layout.get_dim_split_num(1)

            def _bsnd_cp_impl(*args, **kwargs):
                if dsa_cp_fold_enabled():
                    # Head-tail folded query: one call per block, each on its own causal prefix.
                    return fold_lightning_indexer(func, split_id, seq_shards, *args, **kwargs)
                local_q, local_k = args[0], args[1]
                sliced_k = _adjust_bsnd_key(local_k, local_q.shape[1], split_id)
                return func(local_q, sliced_k, *args[2:], **kwargs)
336
337
338
339
340
341
342
343
344
345
346

            if qlen_tensor is None or klen_tensor is None:
                return func(*args, **kwargs)

            if dsa_cp_fold_enabled():
                # See fold_sparse_flash_attention: global cumulative lengths in, per-block out.
                return fold_lightning_indexer(
                    func, seq_shard_id, seq_shards, *args, fold_layout='TND', **kwargs)

            adj_q, adj_k = _adjust_tnd_seq_lens(
                local_q, local_k, qlen_tensor, klen_tensor,
hyper_parallel/core/shard/ops/parallel_npu_dense_lightning_indexer_grad_kl_loss.py
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
            split_id = q_layout.get_split_id(1)
            seq_shards = q_layout.get_dim_split_num(1)

            def _bsnd_cp_impl(*args, **kwargs):
                if dsa_cp_fold_enabled():
                    if len(args) < DENSE_KL_MIN_ARGS["BSND"] or kwargs:
                        raise NotImplementedError(
                            "DSA CP head-tail fold is implemented for the MindSpore positional "
                            "signature of the dense indexer KL loss only.")
                    # d_key_index comes back full-length on the folded layout, so the Partial
                    # reduction and the local narrow in DSAIndexerLossContextParallel still apply.
                    return fold_dense_indexer_kl_loss(func, split_id, seq_shards, *args)
                local_q = args[0]
                s1_local = local_q.shape[1]
                s2_full = args[3].shape[1]
                sliced_k = _adjust_bsnd_key(args[1], s1_local, split_id)
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488

            if qlen_tensor is None or klen_tensor is None:
                return func(*args, **kwargs)

            if dsa_cp_fold_enabled():
                # Global cumulative lengths in, per-block ones out; must precede
                # ``_adjust_tnd_seq_lens``, whose contiguous-slice assumption the fold breaks.
                if len(args) < DENSE_KL_MIN_ARGS["TND"] or kwargs:
                    raise NotImplementedError(
                        "DSA CP head-tail fold is implemented for the MindSpore positional "
                        "signature of the dense indexer KL loss only.")
                return fold_dense_indexer_kl_loss(
                    func, seq_shard_id, seq_shards, *args, fold_layout='TND')

            adj_q, adj_k = _adjust_tnd_seq_lens(
                local_q, local_k, qlen_tensor, klen_tensor,
hyper_parallel/core/shard/ops/parallel_npu_dense_lightning_indexer_softmax_lse.py
383
384
385
386
387
388
389
390
391
392
393
394
            split_id = q_layout.get_split_id(1)
            seq_shards = q_layout.get_dim_split_num(1)

            def _bsnd_cp_impl(*args, **kwargs):
                if dsa_cp_fold_enabled():
                    # Folded layout: the local query holds this rank's two head-tail blocks,
                    # each of which gets its own causal key prefix out of the full-length key.
                    return fold_dense_lightning_indexer_softmax_lse(
                        func, split_id, seq_shards, *args, **kwargs)
                local_q, local_k = args[0], args[1]
                sliced_k = _adjust_bsnd_key(local_k, local_q.shape[1], split_id)
                return func(local_q, sliced_k, *args[2:], **kwargs)
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431

            if qlen_tensor is None or klen_tensor is None:
                return func(*args, **kwargs)

            if dsa_cp_fold_enabled():
                # Folded layout: one call per block against its own causal prefix. The
                # cumulative lengths handed down must be the GLOBAL ones -- restating them
                # per block is the fold's job, and ``_adjust_tnd_seq_lens`` below assumes the
                # contiguous CP slice the fold deliberately breaks -- so this comes first.
                if len(args) <= 4:
                    raise NotImplementedError(
                        "DSA CP head-tail fold is implemented for the MindSpore positional "
                        "signature of the dense indexer forward only.")
                return fold_dense_lightning_indexer_softmax_lse(
                    func, seq_shard_id, seq_shards, *args, fold_layout='TND', **kwargs)

            adj_q, adj_k = _adjust_tnd_seq_lens(
                local_q, local_k, qlen_tensor, klen_tensor,
hyper_parallel/core/shard/ops/parallel_npu_flash_attention_score.py
80
81
82
83
84
85
86
87
88
    is this rank's contiguous CP slice -- so anything keyed on ``cp_rank`` (MindFormers' TND
    softmax-statistics converter repacks per document from ``offset = t * cp_rank``) is wrong
    while this is active and must key on ``split_id`` instead. ``split_num`` is ``2N``.
    """
    return _get_lb_override()


def _normalize_npu_fusion_attention_args(
    query, key, value, head_num, input_layout,
hyper_parallel/core/shard/ops/parallel_npu_sparse_flash_attention.py
409
410
411
412
413
414
415
416
417
418
419
            split_id = q_layout.get_split_id(1)
            seq_shards = q_layout.get_dim_split_num(1)

            def _bsnd_cp_impl(*args, **kwargs):
                if dsa_cp_fold_enabled():
                    # Head-tail folded query: one call per block, each on its own causal prefix.
                    return fold_sparse_flash_attention(func, split_id, seq_shards, *args, **kwargs)
                local_q, local_k, local_v = args[0], args[1], args[2]
                s1_local = local_q.shape[1]
                sliced_k = _adjust_bsnd_key(local_k, s1_local, split_id)
                sliced_v = _adjust_bsnd_key(local_v, s1_local, split_id)
438
439
440
441
442
443
444
445
446
447
448
449
450
451
            qlen_tensor = kwargs.get('actual_seq_lengths_query')
            klen_tensor = kwargs.get('actual_seq_lengths_kv')
            if qlen_tensor is None or klen_tensor is None:
                return func(*args, **kwargs)
            if dsa_cp_fold_enabled():
                # Head-tail folded query: one call per block on its own causal prefix, with the
                # cumulative lengths restated in that block's coordinates. The **global** ones
                # are passed through untouched -- _adjust_tnd_seq_lens assumes a contiguous CP
                # slice, which the fold breaks.
                return fold_sparse_flash_attention(
                    func, seq_shard_id, seq_shards, *args, fold_layout='TND', **kwargs)
            adj_q, adj_k = _adjust_tnd_seq_lens(
                local_q, local_k, qlen_tensor, klen_tensor,
                cp_rank=seq_shard_id,
hyper_parallel/core/shard/ops/parallel_npu_sparse_lightning_indexer_grad_kl_loss.py
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
            split_id = q_layout.get_split_id(1)
            seq_shards = q_layout.get_dim_split_num(1)

            def _bsnd_cp_impl(*args, **kwargs):
                if dsa_cp_fold_enabled():
                    if len(args) < SPARSE_KL_MIN_ARGS["BSND"] or kwargs:
                        raise NotImplementedError(
                            "DSA CP head-tail fold is implemented for the MindSpore positional "
                            "signature of the sparse indexer KL loss only.")
                    # d_key_index comes back full-length on the folded layout, so the Partial
                    # reduction and the local narrow in DSAIndexerLossContextParallel still apply.
                    return fold_sparse_indexer_kl_loss(func, split_id, seq_shards, *args)
                local_q = args[0]
                s1_local = local_q.shape[1]
                # args[3] is the full (unsliced) local key_index; S2 is replicated
                # so local shape == global shape and we can recover s2_full here.
470
471
472
473
474
475
476
477
478

            if qlen_tensor is None or klen_tensor is None:
                return func(*args, **kwargs)

            if dsa_cp_fold_enabled():
                # Head-tail folded query: one call per block on its own causal prefix.
                # The fold wrapper rewrites args 11/12 per block, so the **global**
                # cumulative lengths must reach it untouched.
                # `or kwargs`: fold_sparse_indexer_kl_loss takes no **kwargs, so any keyword
477
478
479
480
481
482
483
484
485
486
487
488
489
490
                # cumulative lengths must reach it untouched.
                # `or kwargs`: fold_sparse_indexer_kl_loss takes no **kwargs, so any keyword
                # argument reaching it would be dropped silently. The BSND branch above already
                # guards this way; keep the two in step.
                if len(args) < SPARSE_KL_MIN_ARGS["TND"] or kwargs:
                    raise NotImplementedError(
                        "DSA CP head-tail fold for the TND indexer KL loss is implemented for "
                        "the MindSpore positional form only; this call passed the sequence "
                        "lengths as keyword arguments (torch path).")
                return fold_sparse_indexer_kl_loss(
                    func, seq_shard_id, seq_shards, *args, fold_layout='TND')

            adj_q, adj_k = _adjust_tnd_seq_lens(
                local_q, local_k, qlen_tensor, klen_tensor,
hyper_parallel/platform/mindspore/platform.py
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399

def _p2p_exchange_raw(tensor, peer_rank: int, group):
    """One symmetric exchange: send ``tensor`` to the peer, return what the peer sent."""
    # pylint: disable=C0415
    from mindspore.mint.distributed import P2POp, batch_isend_irecv, isend, irecv
    send_buf = tensor.contiguous()
    recv_buf = mint.empty_like(send_buf)
    handles = batch_isend_irecv([
        P2POp(isend, send_buf, peer_rank, group),
        P2POp(irecv, recv_buf, peer_rank, group),
    ])
    for handle in handles:
        if handle is not None:
            handle.wait()
    return recv_buf


class _MSP2PExchangeFunction(_Function):
    """Symmetric bidirectional P2P; the gradient takes the same exchange back."""
399
400
401
402
403
404
405
406
407
408
409
410
411
412
    """Symmetric bidirectional P2P; the gradient takes the same exchange back."""

    @staticmethod
    def forward(ctx, tensor, peer_rank: int, group):  # pylint: disable=arguments-differ
        ctx.peer_rank, ctx.group = peer_rank, group
        return _p2p_exchange_raw(tensor, peer_rank, group)

    @staticmethod
    def backward(ctx, grad_output):  # pylint: disable=arguments-differ
        return _p2p_exchange_raw(grad_output, ctx.peer_rank, ctx.group), None, None


class _MSAsyncA2ALazyBwd(_Function):
    """Async all-to-all whose forward and backward both return
1276
1277
1278
1279
1280
1281
1282
1283
1284
1285
1286
        ``batch_isend_irecv`` so the send and the receive overlap on the duplex link.
        """
        # ``peer_rank`` is a GLOBAL rank (that is what ``P2POp`` takes), so compare it with
        # the global rank -- ``dist.get_rank(group)`` would be the rank *within* the group.
        if peer_rank == get_rank_id():
            return tensor
        return _MSP2PExchangeFunction.apply(tensor, peer_rank, group)

    @staticmethod
    def send_object_list(obj_list, dst=None, group=None):
        # pylint: disable=C0415