Diff Coverage

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

Source File Diff Coverage (%) Missing Lines
hyper_parallel/core/shard/ops/parallel_ms_flash_attention_score.py 18.2% 79,84,89,99-102,1038,1040,1046,1051,1054,1067-1069,1075-1076,1097
hyper_parallel/custom_ops/experimental/experimental_ops.py 50.0% 192
hyper_parallel/platform/mindspore/custom_ops/custom_op_impl.py 30.0% 190-192,196-198,200,211-212,214-217,222
hyper_parallel/platform/mindspore/custom_ops/custom_ops.py 66.7% 42
hyper_parallel/platform/torch/custom_ops/__init__.py 0.0% 53,55,60
hyper_parallel/core/shard/ops/parallel_ms_flash_attention_score.py
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93


def set_tnd_softmax_out(enabled: bool) -> None:
    """Whether TND FlashAttention should emit TND-ordered softmax statistics."""
    _TND_SOFTMAX_OUT["enabled"] = bool(enabled)


def tnd_softmax_out_enabled() -> bool:
    """Whether the TND-ordered softmax statistics path is on."""
    return _TND_SOFTMAX_OUT["enabled"]


def _reject_unsupported_v4_inputs(real_shift, drop_mask, padding_mask, prefix, keep_prob) -> None:
    """Raise when an input the aclnn V4 varlen kernel cannot honour is in use."""
    unsupported = [
        name
        for name, value in (
            ("real_shift", real_shift),
            ("drop_mask", drop_mask),
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
            ("prefix", prefix),
        )
        if value is not None
    ]
    if keep_prob is not None and float(keep_prob) != 1.0:
        unsupported.append(f"keep_prob={keep_prob}")
    if unsupported:
        raise NotImplementedError(
            "TND-ordered softmax statistics use the aclnn V4 varlen kernel, which does not "
            f"accept {', '.join(unsupported)}. Turn the path off with "
            "set_tnd_softmax_out(False) for this model."
        )
1034
1035
1036
1037
1038
1039
1040
1041
1042
1043
1044
            seq_split_num = split_info["seq"]

            lb_split_id, lb_split_num = _get_lb_override()

            def _local_kernel(q, k, v, a_qlen, a_kvlen, h_num, s_mode, pre_t, next_t):
                """Run the local FA kernel: stock primitive, or aclnn V4 for TND statistics."""
                if input_layout == "TND" and tnd_softmax_out_enabled():
                    # The V4 operator takes no pse / dropout / padding / prefix inputs and
                    # always runs at keep_prob == 1.0. Refuse rather than forward a subset:
                    # dropping any of them would change the attention numerically while the
                    # output still looks plausible, so it would only ever surface as an
1042
1043
1044
1045
1046
1047
1048
1049
1050
1051
1052
1053
1054
1055
1056
1057
1058
                    # always runs at keep_prob == 1.0. Refuse rather than forward a subset:
                    # dropping any of them would change the attention numerically while the
                    # output still looks plausible, so it would only ever surface as an
                    # accuracy drift, never as an error.
                    _reject_unsupported_v4_inputs(
                        real_shift, drop_mask, padding_mask, prefix, keep_prob
                    )
                    # Imported locally so that processes which are not on TND, or have
                    # this path disabled, never load the custom-operator extension.
                    from hyper_parallel.custom_ops.experimental import (  # pylint: disable=C0415
                        npu_flash_attention_varlen_v4,
                    )
                    s_max, s_sum, attn = npu_flash_attention_varlen_v4(
                        q, k, v, atten_mask=attn_mask,
                        actual_seq_qlen=a_qlen, actual_seq_kvlen=a_kvlen,
                        scale_value=scale_value, head_num=int(h_num), sparse_mode=int(s_mode),
                        pre_tokens=int(pre_t), next_tokens=int(next_t), inner_precise=inner_precise,
1063
1064
1065
1066
1067
1068
1069
1070
1071
1072
1073
                    # placeholder: on TND the stock primitive was measured to return a
                    # shape-(1,) tensor of the same dtype, and that must be reproduced
                    # here. Passing None makes the DTensor dispatch raise
                    # "local_tensor must be a Tensor" when it wraps all four outputs.
                    softmax_out = platform.zeros((1,), dtype=attn.dtype, device=attn.device)
                    return (s_max, s_sum, softmax_out, attn)
                return func(
                    q, k, v, real_shift, drop_mask, padding_mask, attn_mask, prefix,
                    a_qlen, a_kvlen, h_num, keep_prob, scale_value,
                    pre_t, next_t, inner_precise, p_input_layout, s_mode,
                )
1071
1072
1073
1074
1075
1076
1077
1078
1079
1080
                    a_qlen, a_kvlen, h_num, keep_prob, scale_value,
                    pre_t, next_t, inner_precise, p_input_layout, s_mode,
                )

            if head_split_num == 1 and seq_split_num == 1 and lb_split_id is None:
                result = _local_kernel(query, key, value, actual_seq_qlen, actual_seq_kvlen,
                                       head_num, sparse_mode, pre_tokens, next_tokens)
                return FlashAttentionScoreDistributedOp._truncate_result(result)

            adjusted_head_num = self._adjust_head_num(head_num, head_split_num)
1093
1094
1095
1096
1097
1098
1099
1100
1101
                    lb_split_num=lb_split_num,
                ),
            )

            result = _local_kernel(
                query, key, value, adjusted_actual_seq_qlen, adjusted_actual_seq_kvlen,
                int(adjusted_head_num), int(adjusted_sparse_mode),
                int(adjusted_pre_tokens), int(adjusted_next_tokens),
            )
hyper_parallel/custom_ops/experimental/experimental_ops.py
188
189
190
191
192
193
194
195
196

    Returns:
        tuple[Tensor, Tensor, Tensor]: ``(softmax_max, softmax_sum, attention_out)``.
    """
    return _platform.custom_ops.npu_flash_attention_varlen_v4(
        query, key, value, atten_mask, actual_seq_qlen, actual_seq_kvlen,
        scale_value, head_num, sparse_mode, pre_tokens, next_tokens, inner_precise,
        tnd_softmax_out,
    )
hyper_parallel/platform/mindspore/custom_ops/custom_op_impl.py
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
    def forward(ctx, query, key, value, atten_mask, actual_seq_qlen, actual_seq_kvlen,
                scale_value, head_num, sparse_mode, pre_tokens, next_tokens, inner_precise,
                tnd_softmax_out):
        """Forward: returns ``(softmax_max, softmax_sum, attention_out)``."""
        q_len = _to_list_int64(actual_seq_qlen)
        kv_len = _to_list_int64(actual_seq_kvlen)
        outs = _custom_ops.npu_flash_attention_varlen_v4(
            query, key, value, atten_mask, q_len, kv_len, scale_value, head_num,
            sparse_mode, pre_tokens, next_tokens, inner_precise, tnd_softmax_out,
        )
        softmax_max, softmax_sum, attention_out = outs
        ctx.save_for_backward(query, key, value, atten_mask, softmax_max, softmax_sum, attention_out)
        ctx.fa_args = (q_len, kv_len, scale_value, head_num, sparse_mode, pre_tokens, next_tokens,
                       inner_precise, tnd_softmax_out)
        return softmax_max, softmax_sum, attention_out

    @staticmethod
    def backward(ctx, *grad_outputs):
        """Backward through the attention output only.
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
        gradients for the indexer parameters and routes nothing back here -- so their
        incoming gradients are ignored, mirroring what the stock FlashAttentionScore bprop
        does with them.
        """
        query, key, value, atten_mask, softmax_max, softmax_sum, attention_out = ctx.saved_tensors
        (q_len, kv_len, scale_value, head_num, sparse_mode, pre_tokens, next_tokens,
         inner_precise, tnd_softmax) = ctx.fa_args
        dy = grad_outputs[2]
        if dy is None:
            return (None,) * 13
        dq, dk, dv = _custom_ops.npu_flash_attention_varlen_grad_v4(
            query, key, value, dy, atten_mask, softmax_max, softmax_sum, attention_out,
            q_len, kv_len, scale_value, head_num, sparse_mode, pre_tokens, next_tokens,
            inner_precise, tnd_softmax,
        )
        return dq, dk, dv, None, None, None, None, None, None, None, None, None, None


class NpuDenseLightningIndexerGradKlLossDFunction(DFunction):  # pylint: disable=W0221
    """DFunction wrapper for npu_dense_lightning_indexer_grad_kl_loss on MindSpore.
hyper_parallel/platform/mindspore/custom_ops/custom_ops.py
38
39
40
41
42
43
44
45
46

    @staticmethod
    def npu_flash_attention_varlen_v4(*args, **kwargs):
        """TND varlen FlashAttention via aclnn V4 (differentiable)."""
        return NpuFlashAttentionVarLenV4DFunction.apply(*args, **kwargs)

    @staticmethod
    def npu_dense_lightning_indexer_softmax_lse(*args, **kwargs):
        """Compute dense lightning indexer softmax log-sum-exp via custom NPU operator."""
hyper_parallel/platform/torch/custom_ops/__init__.py
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
            "on the PyTorch platform."
        )

    @staticmethod
    def npu_flash_attention_varlen_v4(*args, **kwargs):
        """TND varlen FlashAttention via aclnn V4; not supported on PyTorch."""
        raise NotImplementedError(
            "npu_flash_attention_varlen_v4 is MindSpore-only: it exists to reach aclnn V4's "
            "softmaxOutLayout, which torch_npu's FA wrapper does not expose."
        )

    @staticmethod
    def npu_mhc_post(*args, **kwargs):
        """NPU MHC post-processing operator; not supported on PyTorch."""
        raise NotImplementedError(
            "npu_mhc_post is not supported on the PyTorch platform."