已开启
【RFC】DeepSeek-V4 CSA/HCA 自定义算子接入与对外接口对标 torch_extension #244
hedongdong创建于  6月24日
hedongdong成员
6月24日 创建

【RFC】DeepSeek-V4 CSA/HCA 自定义算子接入与对外接口对标 torch_extension

背景

DeepSeek-V4 在 DSA 动态稀疏注意力(见 #150)基础上引入 CSA(Compressed Sparse Attention)/ HCA

  1. Lightning Indexer 压缩 —— Indexer 支持对 KV 做压缩(cmp_ratio = 1 / 4 / 128),以 cmp_residual_k(原始 k 长度 % cmp_ratio)描述压缩有效范围。
  2. MLA sparse attention —— 以 sparse_flash_mla triplet 替代原 shared-KV 三件套,支持 ori_kv(band/滑窗)+ cmp_kv(压缩)双分支与 attention-sink。

需接入对应的 6 个 Ascend NPU 自定义算子,并保证三个对外接口严格对标 ops-transformer torch_extensionschema()

目标:算子接入

对外接口对标 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) 命名:topksparse_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_modesparse_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) 同(qquery 4 个位置参数一致
ori_kvcmp_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

likedislike
Hhedongdong成员
7月4日 修改了issue 的描述
Hhedongdong成员
7月15日 关联了pull request:feat: add DeepSeek-V4 CSA/HCA custom ops (MLA sparse attention + lightning indexer)