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 82.5% 163,176,214-216,227,257-264,266-272,274-280,282-284,290,293,323,356-357,359,407,456,494,497,502,539,542-543,549,637,640,687,697,700-701,710
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
159
160
161
162
163
164
165
166
167
    if idx is None or type(idx).__module__.split(".")[0] != backend:
        idx = platform.from_numpy(np.asarray(values, dtype=np.int32))
        _INDEX_CACHE[key] = idx
    if backend == "torch":
        idx = idx.to(ref.device)
    return idx


def _chunked_shape(shape: tuple, seq_dim: int, chunks: int) -> tuple:
172
173
174
175
176
177
178
179
180
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)
210
211
212
213
214
215
216
217
218
219
220
        ``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.
223
224
225
226
227
228
229
230
231
    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):
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
282
283
284
285
286
287
288
    @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."""
286
287
288
289
290
291
292
293
294
295
296

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)
319
320
321
322
323
324
325
326

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)

352
353
354
355
356
357
358
359
360
361
362
363
        return
    _KLEN_CHECKED[0] = True
    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. "
403
404
405
406
407
408
409
410
411
    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
452
453
454
455
456
457
458
459
460
    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(
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506

    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,
535
536
537
538
539
540
541
542
543
544
545
546
547
            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)
545
546
547
548
549
550
551
552

    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)
633
634
635
636
637
638
639
640
641
642
643
644
    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)
683
684
685
686
687
688
689
690
691
    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)
693
694
695
696
697
698
699
700
701
702
703
704
705

    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)
706
707
708
709
710
711
712
713
714
    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:
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