mindspore-ai / mindspore-ai/hyper-parallel
【RFC】DeepSeek-V4 CSA/HCA 自定义算子接入与对外接口对标 torch_extension
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 53
- Forks
- 63
- Avg merge
- 23h 45m
- Merged PRs (30d)
- 63
Description
【RFC】DeepSeek-V4 CSA/HCA 自定义算子接入与对外接口对标 torch_extension
背景
DeepSeek-V4 在 DSA 动态稀疏注意力(见 #150)基础上引入 CSA(Compressed Sparse Attention)/ HCA:
- Lightning Indexer 压缩 —— Indexer 支持对 KV 做压缩(cmp_ratio = 1 / 4 / 128),以
cmp_residual_k(原始 k 长度 % cmp_ratio)描述压缩有效范围。 - MLA sparse attention —— 以
sparse_flash_mlatriplet 替代原 shared-KV 三件套,支持 ori_kv(band/滑窗)+ cmp_kv(压缩)双分支与 attention-sink。
需接入对应的 6 个 Ascend NPU 自定义算子,并保证三个对外接口严格对标 ops-transformer torch_extension 的 schema()。
目标:算子接入
-
npu_lightning_indexer(底层lightning_indexer)—— 稀疏 token 索引,cmp_ratio 1/4/128 -
npu_sparse_flash_mla(底层sparse_flash_mla)—— MLA 稀疏注意力前向 -
npu_sparse_flash_mla_grad(底层sparse_flash_mla_grad)—— MLA 稀疏注意力反向;DFunction.backward 内部调起,同时暴露为对外接口供网络自定义反向消费softmax_l1_norm -
npu_sparse_lightning_indexer_kl_loss_grad(底层..._kl_loss_grad)—— Indexer 反向 + KL 损失 - tiling metadata —— 由各主算子在 kernel cc 内部从张量 shape 直接内联计算,不再作为独立自定义算子
对外接口对标 torch_extension
对标原则:① *(位置 vs 强制 kwarg)分隔点与标杆一致,位置参数个数相同;② 具名入参相对顺序与标杆一致;③ 两份 layout 合并为单一 layout(默认 BSND),DFunction 内部双传 layout_q/layout_k;④ 对外接口不出现标杆不存在的入参(仅允许命名差异,逐项标注)。
1. npu_lightning_indexer
标杆 lightning_indexer:
lightning_indexer(Tensor q, Tensor k, Tensor w, int topk, *,
Tensor? cu_seqlens_q, Tensor? cu_seqlens_k, Tensor? seqused_q, Tensor? seqused_k,
Tensor? cmp_residual_k, Tensor? block_table, Tensor? output_idx_offset, Tensor? metadata,
int max_seqlen_q=-1, str layout_q="BSND", str layout_k="BSND",
int mask_mode=0, int cmp_ratio=1, int return_value=0) -> (Tensor, Tensor)
hyper-parallel:
npu_lightning_indexer(query, key, weights, sparse_count, *,
cu_seq_lens_q=None, cu_seq_lens_k=None, cmp_residual_k=None, block_table=None,
layout='BSND', sparse_mode=0, cmp_ratio=1, return_value=False)
| # | 标杆 | 默认 | hyper-parallel | 差异 |
|---|---|---|---|---|
| 1-3 | q / k / w (pos) |
— | query / key / weights (pos) |
语义名 |
| 4 | topk (pos) |
— | sparse_count (pos) |
命名:topk↔sparse_count |
| 5 | cu_seqlens_q |
None | cu_seq_lens_q |
命名缩写(累积前缀和语义) |
| 6 | cu_seqlens_k |
None | cu_seq_lens_k |
同上 |
| 7 | seqused_q |
None | — | 缺;对照实验 lightning golden 不引用,不参与 |
| 8 | seqused_k |
None | — | 缺;同上 |
| 9 | cmp_residual_k |
None | cmp_residual_k |
已暴露:实测参与计算(residual 0 vs 2 输出不同) |
| 10 | block_table |
None | block_table |
一致 |
| 11 | output_idx_offset |
None | — | 缺;对照实验不参与 |
| 12 | metadata |
None | — | 缺;DFunction 内部固定 None |
| 13 | max_seqlen_q |
-1 | — | 缺;DFunction 内部固定 -1(自动推导) |
| 14-15 | layout_q / layout_k |
"BSND" | layout |
合并为单一 layout |
| 16 | mask_mode |
0 | sparse_mode |
命名:mask_mode↔sparse_mode |
| 17 | cmp_ratio |
1 | cmp_ratio |
一致 |
| 18 | return_value |
0 | return_value |
int↔bool,语义一致 |
2. npu_sparse_flash_mla
标杆 sparse_flash_mla(主算子,仅 1 个位置参数 q):
sparse_flash_mla(Tensor q, *, ori_kv, cmp_kv, ori_sparse_indices, cmp_sparse_indices,
ori_block_table, cmp_block_table, cu_seqlens_q, cu_seqlens_ori_kv, cu_seqlens_cmp_kv,
seqused_q, seqused_ori_kv, seqused_cmp_kv, cmp_residual_kv, ori_topk_length, cmp_topk_length,
sinks, metadata, float softmax_scale=1.0, int cmp_ratio=1, int ori_mask_mode=4,
int cmp_mask_mode=3, int ori_win_left=127, int ori_win_right=0,
str layout_q="BSND", str layout_kv="BSND", int topk_value_mode=1, bool return_softmax_lse=False)
hyper-parallel:
npu_sparse_flash_mla(query, *, ori_kv=None, cmp_kv=None, cmp_sparse_indices=None,
cu_seq_lens_q=None, cu_seq_lens_ori_kv=None, cu_seq_lens_cmp_kv=None,
seqused_q=None, seqused_ori_kv=None, seqused_cmp_kv=None, cmp_residual_kv=None, sinks=None,
softmax_scale=1.0, cmp_ratio=1, ori_mask_mode=4, cmp_mask_mode=3,
ori_win_left=127, ori_win_right=0, layout='BSND', return_softmax_lse=False)
| # | 标杆 | hyper-parallel | 差异 |
|---|---|---|---|
q (pos) |
query (pos) |
位置参数仅 1 个(与标杆一致) | |
ori_kv / cmp_kv |
ori_kv / cmp_kv |
一致 | |
ori_sparse_indices |
— | 缺;对照实验确认 BSND/TND 传入 vs None 输出零变化(kernel 忽略 + golden 无条件置 None),不参与 | |
cmp_sparse_indices |
cmp_sparse_indices |
一致 | |
ori_block_table / cmp_block_table |
— | 缺;PA 专用,非 PA 连续布局不涉及 | |
cu_seqlens_q/ori_kv/cmp_kv |
cu_seq_lens_q/ori_kv/cmp_kv |
命名缩写 | |
seqused_q |
seqused_q |
已暴露:BSND 声明有效 query 行数(不改有效行数值,无效行未初始化) | |
seqused_ori_kv / seqused_cmp_kv |
seqused_ori_kv / seqused_cmp_kv |
一致 | |
cmp_residual_kv |
cmp_residual_kv |
一致;CANN 在 cmp_ratio≠1 且 cmp_mask_mode=3 时强制要求 | |
ori_topk_length / cmp_topk_length |
— | 缺;与 full 模式互斥 | |
sinks |
sinks |
一致 | |
metadata |
— | 缺;kernel 内部自动计算并消费 | |
softmax_scale=1.0 / cmp_ratio=1 / ori_mask_mode=4 / cmp_mask_mode=3 / ori_win_left=127 / ori_win_right=0 |
同名同默认 | 一致 | |
layout_q / layout_kv |
layout |
合并 | |
topk_value_mode=1 |
— | 缺;DFunction 内部固定 1(对照实验不参与) | |
return_softmax_lse |
return_softmax_lse |
一致 |
3. npu_sparse_lightning_indexer_kl_loss_grad
标杆(5 个位置参数):
sparse_lightning_indexer_kl_loss_grad(q, k, w, sparse_indices, attn_softmax_l1_norm, *,
cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k, cmp_residual_k, metadata,
str layout_q="TND", str layout_k="TND", int mask_mode=3, int cmp_ratio=1)
hyper-parallel:
npu_sparse_lightning_indexer_kl_loss_grad(query, key, weights, sparse_indices, attn_softmax_l1_norm, *,
cu_seq_lens_q=None, cu_seq_lens_k=None, seqused_q=None, seqused_k=None, cmp_residual_k=None,
layout='BSND', mask_mode=3, cmp_ratio=1)
| # | 标杆 | 默认 | hyper-parallel | 默认 | 差异 |
|---|---|---|---|---|---|
| 1-5 | q/k/w/sparse_indices/attn_softmax_l1_norm (pos) |
— | 同(语义名) | — | 5 个位置参数一致 |
| 6-7 | cu_seqlens_q/k |
None | cu_seq_lens_q/k |
None | 命名缩写 |
| 8-9 | seqused_q/k |
None | seqused_q/k |
None | 一致 |
| 10 | cmp_residual_k |
None | cmp_residual_k |
None | 一致 |
| 11 | metadata |
None | — | — | 缺;kernel 内部自动计算 |
| 12-13 | layout_q / layout_k |
"TND" | layout |
"BSND" | 合并;默认值不同:对外统一默认 BSND(BSND/TND 均验证对齐) |
| 14 | mask_mode |
3 | mask_mode |
3 | 一致 |
| 15 | cmp_ratio |
1 | cmp_ratio |
1 | 一致 |
4. npu_sparse_flash_mla_grad
主算子 npu_sparse_flash_mla 的反向:DFunction.backward 内部调起,同时暴露为对外接口,供网络在自定义反向内获取 softmax_l1_norm(主注意力目标分布 p),直接喂给 npu_sparse_lightning_indexer_kl_loss_grad。
标杆 sparse_flash_mla_grad(4 个位置参数):
sparse_flash_mla_grad(q, dout, attn_out, softmax_lse, *, ori_kv, cmp_kv,
ori_sparse_indices, cmp_sparse_indices, cu_seqlens_q, cu_seqlens_ori_kv, cu_seqlens_cmp_kv,
seqused_q, seqused_ori_kv, seqused_cmp_kv, cmp_residual_kv, ori_topk_length, cmp_topk_length,
sinks, metadata, float softmax_scale, int cmp_ratio, int ori_mask_mode, int cmp_mask_mode,
int ori_win_left, int ori_win_right, str layout_q="BSND", str layout_kv="BSND")
hyper-parallel:
npu_sparse_flash_mla_grad(query, dout, attn_out, softmax_lse, *, ori_kv=None, cmp_kv=None,
ori_sparse_indices=None, cmp_sparse_indices=None,
cu_seq_lens_q=None, cu_seq_lens_ori_kv=None, cu_seq_lens_cmp_kv=None,
seqused_q=None, seqused_ori_kv=None, seqused_cmp_kv=None, cmp_residual_kv=None,
ori_topk_length=None, cmp_topk_length=None, sinks=None,
softmax_scale=1.0, cmp_ratio=1, ori_mask_mode=4, cmp_mask_mode=3,
ori_win_left=127, ori_win_right=0, layout='BSND')
| # | 标杆 | hyper-parallel | 差异 |
|---|---|---|---|
q/dout/attn_out/softmax_lse (pos) |
同(q→query) |
4 个位置参数一致 | |
ori_kv … cmp_topk_length / sinks |
同名 | 一致(cu_seqlens_*↔cu_seq_lens_* 命名缩写) |
|
metadata |
— | 缺;grad kernel 内部自推 tiling(固定 nullptr) | |
softmax_scale / cmp_ratio / ori_mask_mode / cmp_mask_mode / ori_win_left / ori_win_right |
同名同默认 | 一致 | |
layout_q / layout_kv |
layout |
合并 |
返回 (d_query, d_ori_kv, d_cmp_kv, d_sinks, ori_softmax_l1_norm, cmp_softmax_l1_norm):后两者为 reduceG(softmax)/G 的主注意力分布,shape 跟随 ori/cmp_sparse_indices。与 DFunction.backward 调用同一裸 kernel(seqused_* / ori/cmp_topk_length 在 DFunction 内部固定 None,对外暴露为默认 None 的 optional)。
缺失入参的用途分析(对照实验)
"数据生成器不喂某参数"仅是弱证据。对所有非 PA 缺失参数做对照实验(固定其余输入,传非平凡值 vs None,看真实 NPU kernel 输出是否变化):
| 参数 | layout | 对照结果 | 结论 |
|---|---|---|---|
cmp_residual(三算子) |
BSND/TND | 0→2 变(val max≈7.97) | 参与 → 已预留并补用例 |
ori_sparse_indices(flash_mla) |
BSND+TND | 合法索引→None 输出零变化 | 不参与(kernel 忽略 + golden 置 None) |
ori_topk_length / cmp_topk_length |
BSND | 传入即报错 | 与 full 模式互斥 |
topk_value_mode |
BSND | 1→0 不变 | 不参与 |
seqused_q(flash_mla) |
BSND | 仅声明有效行数;不改有效行数值 | 已暴露,向后兼容 |
lightning seqused_q/k / max_seqlen_q / output_idx_offset |
TND | 不变 | 不参与 |
PA 专用的
ori/cmp_block_table按约定不关注(非 PA 连续布局结构上不生成)。
测试
tests/mindspore/st/custom_ops/experimental/:对外接口用例 test_experimental_interfaces.py 覆盖 BSND + TND × cmp_ratio 1 / 4 / 128(含 metadata 内联路径、反向 softmax_l1_norm 内容与标杆对齐),统一 float16,与标杆 benchmark 逐元素对齐(assert_array_equal):23 passed。
schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 244
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/244
Contributor guide
No contributing guide indexed for this repository
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start with tests/mindspore/st/custom_ops/experimental/test_experimental_interfaces.py and compare the listed npu_* schemas with the torch_extension interfaces for BSND and TND. Run the 23 interface tests across cmp_ratio 1, 4, and 128; done means the custom operators and exposed backward outputs align element-for-element with the benchmark.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python, pytorch
- Domain
- backend, machine-learning
- Issue type
- Feature
- Difficulty
- 5/5
- Estimated time
- Over a week
- Activity status
- Active
- Clarity
- Clearly specified
- Newbie friendliness
- 20/100