mindspore-ai / mindspore-ai/hyper-parallel

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

Open
#678 0 comments 0 reactions 0 assignees View on GitHub

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

  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()

目标:算子接入

  • 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) 命名: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

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

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. 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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.