已开启
【RFC】DeepSeek-V4 CSA/HCA 自定义算子接入与对外接口对标 torch_extension #244
hedongdong创建于 6月24日
7月4日 修改了issue 的描述
Hhedongdong
7月15日 关联了pull request:feat: add DeepSeek-V4 CSA/HCA custom ops (MLA sparse attention + lightning indexer)
7月15日 关联了pull request:feat: add DeepSeek-V4 CSA/HCA custom ops (MLA sparse attention + lightning indexer)
【RFC】DeepSeek-V4 CSA/HCA 自定义算子接入与对外接口对标 torch_extension
背景
DeepSeek-V4 在 DSA 动态稀疏注意力(见 #150)基础上引入 CSA(Compressed Sparse Attention)/ HCA:
cmp_residual_k(原始 k 长度 % cmp_ratio)描述压缩有效范围。sparse_flash_mlatriplet 替代原 shared-KV 三件套,支持 ori_kv(band/滑窗)+ cmp_kv(压缩)双分支与 attention-sink。需接入对应的 6 个 Ascend NPU 自定义算子,并保证三个对外接口严格对标 ops-transformer
torch_extension的schema()。目标:算子接入
对外接口对标 torch_extension
对标原则:①
*(位置 vs 强制 kwarg)分隔点与标杆一致,位置参数个数相同;② 具名入参相对顺序与标杆一致;③ 两份 layout 合并为单一layout(默认BSND),DFunction 内部双传 layout_q/layout_k;④ 对外接口不出现标杆不存在的入参(仅允许命名差异,逐项标注)。1.
npu_lightning_indexer标杆
lightning_indexer:hyper-parallel:
q/k/w(pos)query/key/weights(pos)topk(pos)sparse_count(pos)topk↔sparse_countcu_seqlens_qcu_seq_lens_qcu_seqlens_kcu_seq_lens_kseqused_qseqused_kcmp_residual_kcmp_residual_kblock_tableblock_tableoutput_idx_offsetmetadatamax_seqlen_qlayout_q/layout_klayoutmask_modesparse_modemask_mode↔sparse_modecmp_ratiocmp_ratioreturn_valuereturn_value2.
npu_sparse_flash_mla标杆
sparse_flash_mla(主算子,仅 1 个位置参数q):hyper-parallel:
q(pos)query(pos)ori_kv/cmp_kvori_kv/cmp_kvori_sparse_indicescmp_sparse_indicescmp_sparse_indicesori_block_table/cmp_block_tablecu_seqlens_q/ori_kv/cmp_kvcu_seq_lens_q/ori_kv/cmp_kvseqused_qseqused_qseqused_ori_kv/seqused_cmp_kvseqused_ori_kv/seqused_cmp_kvcmp_residual_kvcmp_residual_kvori_topk_length/cmp_topk_lengthsinkssinksmetadatasoftmax_scale=1.0 /cmp_ratio=1 /ori_mask_mode=4 /cmp_mask_mode=3 /ori_win_left=127 /ori_win_right=0layout_q/layout_kvlayouttopk_value_mode=1return_softmax_lsereturn_softmax_lse3.
npu_sparse_lightning_indexer_kl_loss_grad标杆(5 个位置参数):
hyper-parallel:
q/k/w/sparse_indices/attn_softmax_l1_norm(pos)cu_seqlens_q/kcu_seq_lens_q/kseqused_q/kseqused_q/kcmp_residual_kcmp_residual_kmetadatalayout_q/layout_klayoutmask_modemask_modecmp_ratiocmp_ratio4.
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 个位置参数):hyper-parallel:
q/dout/attn_out/softmax_lse(pos)q→query)ori_kv…cmp_topk_length/sinkscu_seqlens_*↔cu_seq_lens_*命名缩写)metadatasoftmax_scale/cmp_ratio/ori_mask_mode/cmp_mask_mode/ori_win_left/ori_win_rightlayout_q/layout_kvlayout返回
(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 输出是否变化):
cmp_residual(三算子)ori_sparse_indices(flash_mla)ori_topk_length/cmp_topk_lengthtopk_value_modeseqused_q(flash_mla)seqused_q/k/max_seqlen_q/output_idx_offset测试
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。