已合并
torch_extension目录整改 #6973
torch_extension目录整改 #6973
已合并
jinying创建于 6月16日
85 个文件变更+1756-1205
@@ -2,8 +2,8 @@ import torch
2import torch_npu2import torch_npu
3import math3import math
4import numpy as np4import numpy as np
5-import npu_ops_transformer5+import cann_ops_transformer
6-from npu_ops_transformer.ops import npu_flash_attn6+from cann_ops_transformer.ops import npu_flash_attn
7torch.manual_seed(42)7torch.manual_seed(42)
8 8 
9# B = 19# B = 1
@@ -15,9 +15,9 @@ import math
15import numpy as np15import numpy as np
16import random16import random
17from einops import rearrange17from einops import rearrange
18-import npu_ops_transformer18+import cann_ops_transformer
19-from npu_ops_transformer.ops import npu_flash_attn19+from cann_ops_transformer.ops import npu_flash_attn
20-from npu_ops_transformer.ops import npu_flash_attn_metadata20+from cann_ops_transformer.ops import npu_flash_attn_metadata
21from utils import trans_bnsd_to_layout21from utils import trans_bnsd_to_layout
22import torchair22import torchair
23from torchair.configs.compiler_config import CompilerConfig23from torchair.configs.compiler_config import CompilerConfig
@@ -106,7 +106,7 @@ def flash_attn_metadata_only(**kwargs):
106 "layout_out": layout_out,106 "layout_out": layout_out,
107 })107 })
108 108 
109- metadata = torch.ops.npu_ops_transformer.npu_flash_attn_metadata(109+ metadata = torch.ops.cann_ops_transformer.npu_flash_attn_metadata(
110 cu_seqlens_q = cu_q_u32,110 cu_seqlens_q = cu_q_u32,
111 cu_seqlens_kv = cu_kv_u32,111 cu_seqlens_kv = cu_kv_u32,
112 seqused_q = seqused_q_u32,112 seqused_q = seqused_q_u32,
@@ -261,7 +261,7 @@ def flash_attn_npu(q, k, v, q_rope, k_rope, atten_mask, **kwargs):
261 "layout_out": layout_out,261 "layout_out": layout_out,
262 })262 })
263 263 
264- metadata = torch.ops.npu_ops_transformer.npu_flash_attn_metadata(264+ metadata = torch.ops.cann_ops_transformer.npu_flash_attn_metadata(
265 cu_seqlens_q = cu_q_meta,265 cu_seqlens_q = cu_q_meta,
266 cu_seqlens_kv = cu_kv_meta,266 cu_seqlens_kv = cu_kv_meta,
267 seqused_q = seqused_q_meta,267 seqused_q = seqused_q_meta,
@@ -483,7 +483,7 @@ def flash_attn_npu_graph(q, k, v, q_rope, k_rope, atten_mask, **kwargs):
483 "layout_out": layout_out,483 "layout_out": layout_out,
484 })484 })
485 485 
486- metadata = torch.ops.npu_ops_transformer.npu_flash_attn_metadata(486+ metadata = torch.ops.cann_ops_transformer.npu_flash_attn_metadata(
487 cu_seqlens_q = cu_q_meta,487 cu_seqlens_q = cu_q_meta,
488 cu_seqlens_kv = cu_kv_meta,488 cu_seqlens_kv = cu_kv_meta,
489 seqused_q = seqused_q_meta,489 seqused_q = seqused_q_meta,
@@ -8,7 +8,7 @@ from typing import Any, Dict
8from .base import Backend8from .base import Backend
9 9 
10try:10try:
11- from npu_ops_transformer.ops import npu_flash_attn, npu_flash_attn_metadata11+ from cann_ops_transformer.ops import npu_flash_attn, npu_flash_attn_metadata
12 _HAS_NPU = True12 _HAS_NPU = True
13except ImportError:13except ImportError:
14 _HAS_NPU = False14 _HAS_NPU = False
@@ -99,7 +99,7 @@ pytests/
99| 依赖 | 说明 |99| 依赖 | 说明 |
100|------|------|100|------|------|
101| `torch_npu` | PyTorch NPU扩展 |101| `torch_npu` | PyTorch NPU扩展 |
102-| `npu_ops_transformer` | 提供`npu_flash_attn`和`npu_flash_attn_metadata` |102+| `cann_ops_transformer` | 提供`npu_flash_attn`和`npu_flash_attn_metadata` |
103 103 
104### GPU模式(可选)104### GPU模式(可选)
105 105 
@@ -112,7 +112,7 @@ pytests/
112 112 
113```bash113```bash
114# NPU114# NPU
115-python -c "import torch_npu, npu_ops_transformer; print('NPU OK')"115+python -c "import torch_npu, cann_ops_transformer; print('NPU OK')"
116 116 
117# GPU117# GPU
118python -c "import torch, flash_attn, einops; print('GPU OK')"118python -c "import torch, flash_attn, einops; print('GPU OK')"
@@ -425,9 +425,9 @@ flash_attn_inputs = (
425 425 
426## 十一、常见问题426## 十一、常见问题
427 427 
428-### Q1: `ModuleNotFoundError: No module named 'npu_ops_transformer'`428+### Q1: `ModuleNotFoundError: No module named 'cann_ops_transformer'`
429 429 
430-NPU模式需要安装`npu_ops_transformer`,或改用GPU模式:430+NPU模式需要安装`cann_ops_transformer`,或改用GPU模式:
431```bash431```bash
432python test_flash_attn.py --case_id BASE_01 --use_gpu432python test_flash_attn.py --case_id BASE_01 --use_gpu
433```433```
@@ -1,7 +1,7 @@
1import subprocess 1import subprocess
2import torch 2import torch
3import torch_npu 3import torch_npu
4-import npu_ops_transformer 4+import cann_ops_transformer
5 5
6# 初始化 NPU 6# 初始化 NPU
7torch_npu.npu.set_device(0) 7torch_npu.npu.set_device(0)
@@ -40,8 +40,8 @@ print("seqused_kv",seqused_kv)
40print("seqused_kv",seqused_kv)40print("seqused_kv",seqused_kv)
41 41
42# 调用算子42# 调用算子
43-# result = npu_ops_transformer.ops.npu_flash_attn_metadata(43+# result = cann_ops_transformer.ops.npu_flash_attn_metadata(
44-result = torch.ops.npu_ops_transformer.npu_flash_attn_metadata(44+result = torch.ops.cann_ops_transformer.npu_flash_attn_metadata(
45 # cu_seqlens_q = cu_seqlens_q,45 # cu_seqlens_q = cu_seqlens_q,
46 # cu_seqlens_kv = cu_seqlens_kv,46 # cu_seqlens_kv = cu_seqlens_kv,
47 # seqused_q = seqused_q,47 # seqused_q = seqused_q,
@@ -21,7 +21,7 @@ def build_case_script(case: Dict, runs: int, case_name: str, device_id: int = 0)
21 "import torch, torch_npu",21 "import torch, torch_npu",
22 f"torch.npu.set_device({device_id})",22 f"torch.npu.set_device({device_id})",
23 "from core.data import flash_attn_inputs",23 "from core.data import flash_attn_inputs",
24- "from npu_ops_transformer.ops import npu_flash_attn, npu_flash_attn_metadata",24+ "from cann_ops_transformer.ops import npu_flash_attn, npu_flash_attn_metadata",
25 "device = torch.device('npu:0')",25 "device = torch.device('npu:0')",
26 f"c = {params_repr}",26 f"c = {params_repr}",
27 f"print('CASE {case_name} {mode} runs={runs}')",27 f"print('CASE {case_name} {mode} runs={runs}')",
@@ -315,7 +315,7 @@ def build_batch_script(cases: List[Dict], runs: int, cold_thr: int,
315 "import torch, torch_npu",315 "import torch, torch_npu",
316 f"torch.npu.set_device({device_id})",316 f"torch.npu.set_device({device_id})",
317 "from core.data import flash_attn_inputs",317 "from core.data import flash_attn_inputs",
318- "from npu_ops_transformer.ops import npu_flash_attn, npu_flash_attn_metadata",318+ "from cann_ops_transformer.ops import npu_flash_attn, npu_flash_attn_metadata",
319 "device = torch.device('npu:0')",319 "device = torch.device('npu:0')",
320 f"_hot_cases = {case_dicts['hot']}",320 f"_hot_cases = {case_dicts['hot']}",
321 f"_cold_cases = {case_dicts['cold']}",321 f"_cold_cases = {case_dicts['cold']}",
@@ -19,8 +19,8 @@ import numpy as np
19import math19import math
20import ctypes20import ctypes
21import copy21import copy
22-import npu_ops_transformer22+import cann_ops_transformer
23-from npu_ops_transformer.ops import npu_lightning_indexer_v223+from cann_ops_transformer.ops import lightning_indexer_v2
24 24 
25class GeneralizedLIV2:25class GeneralizedLIV2:
26 def __init__(self, batch_size, q_seq, k_seq, q_t_size, k_t_size, q_head_num, k_head_num,26 def __init__(self, batch_size, q_seq, k_seq, q_t_size, k_t_size, q_head_num, k_head_num,
@@ -606,7 +606,7 @@ def liv2_output_single(params):
606 block_table = torch.from_numpy(block_table).to(dtype=torch.int32).npu()606 block_table = torch.from_numpy(block_table).to(dtype=torch.int32).npu()
607 max_seqlen_q = actual_seq_lengths_query.max().item()607 max_seqlen_q = actual_seq_lengths_query.max().item()
608 max_seqlen_k = actual_seq_lengths_key.max().item()608 max_seqlen_k = actual_seq_lengths_key.max().item()
609- # metadata = torch.ops.custom.npu_lightning_indexer_v2_metadata(609+ # metadata = torch.ops.custom.lightning_indexer_v2_metadata(
610 # num_heads_q = q_head_num,610 # num_heads_q = q_head_num,
611 # num_heads_k = k_head_num,611 # num_heads_k = k_head_num,
612 # head_dim = head_dim,612 # head_dim = head_dim,
@@ -628,7 +628,7 @@ def liv2_output_single(params):
628 628 
629 # metadata = metadata.npu()629 # metadata = metadata.npu()
630 metadata = None630 metadata = None
631- npu_result, _ = npu_lightning_indexer_v2(query, key, weights,631+ npu_result, _ = lightning_indexer_v2(query, key, weights,
632 cu_seqlens_q = cu_seqlens_q,632 cu_seqlens_q = cu_seqlens_q,
633 cu_seqlens_k = cu_seqlens_k,633 cu_seqlens_k = cu_seqlens_k,
634 seqused_q = seqused_q,634 seqused_q = seqused_q,
@@ -42,7 +42,7 @@
42|ori_topk_length|可选输入|预留参数,当前不生效|INT32|ND|42|ori_topk_length|可选输入|预留参数,当前不生效|INT32|ND|
43|cmp_topk_length|可选输入|预留参数,当前不生效|INT32|ND|43|cmp_topk_length|可选输入|预留参数,当前不生效|INT32|ND|
44|sinks|可选输入|注意力下沉tensor|FLOAT32|ND|44|sinks|可选输入|注意力下沉tensor|FLOAT32|ND|
45-|metadata|可选输入|aicpu算子(npu_mixed_quant_sparse_flash_mla_metadata)的分核结果,shape固定为[1024]|INT32|ND|45+|metadata|可选输入|aicpu算子(mixed_quant_sparse_flash_mla_metadata)的分核结果,shape固定为[1024]|INT32|ND|
46|quant_mode|可选属性|默认值为None,表示K、V nope的量化模式,当前仅支持1、2,1表示K、V nope为per_token_group量化,scale类型为bfloat16,2表示K、V nope为per_token_group量化,scale类型为float8_e8m0|INT32|-|46|quant_mode|可选属性|默认值为None,表示K、V nope的量化模式,当前仅支持1、2,1表示K、V nope为per_token_group量化,scale类型为bfloat16,2表示K、V nope为per_token_group量化,scale类型为float8_e8m0|INT32|-|
47|rope_head_dim|可选属性|默认值为None,当前仅支持64|INT32|-|47|rope_head_dim|可选属性|默认值为None,当前仅支持64|INT32|-|
48|softmax_scale|可选属性|默认值为None,代表缩放系数,作为q与ori_kv和cmp_kv矩阵乘后Muls的scalar值|FLOAT32|-|48|softmax_scale|可选属性|默认值为None,代表缩放系数,作为q与ori_kv和cmp_kv矩阵乘后Muls的scalar值|FLOAT32|-|
@@ -79,5 +79,5 @@
79- Q\_S和S1表示q shape中的S,S2表示ori\_kv shape中的S,Q\_N和N1表示num\_q\_heads,KV\_N和N2表示num\_ori\_kv\_heads和num\_cmp\_kv\_heads;T1表示q shape中的T。79- Q\_S和S1表示q shape中的S,S2表示ori\_kv shape中的S,Q\_N和N1表示num\_q\_heads,KV\_N和N2表示num\_ori\_kv\_heads和num\_cmp\_kv\_heads;T1表示q shape中的T。
80 80 
81## 调用说明81## 调用说明
82-- 调用方式:使用npu_ops_tranformer包中的npu_mixed_quant_sparse_flash_mla接口进行调用,82+- 调用方式:使用npu_ops_tranformer包中的mixed_quant_sparse_flash_mla接口进行调用,
83- 详见torch_extension/npu_ops_transformer/ops/mixed_quant_sparse_flash_mla.py83+ 详见torch_extension/cann_ops_transformer/ops/mixed_quant_sparse_flash_mla.py
@@ -18,7 +18,7 @@ import pytest
18import random18import random
19import numpy as np19import numpy as np
20import math20import math
21-import npu_ops_transformer21+import cann_ops_transformer
22import custom_ops as ops22import custom_ops as ops
23import torchair23import torchair
24from torchair.configs.compiler_config import CompilerConfig24from torchair.configs.compiler_config import CompilerConfig
@@ -31,7 +31,7 @@ class Network(torch.nn.Module):
31 cmp_block_table, cu_seqlens_q, seqused_q, seqused_ori_kv, seqused_cmp_kv, cmp_residual_kv, sinks, metadata, kv_quant_mode, rope_head_dim,31 cmp_block_table, cu_seqlens_q, seqused_q, seqused_ori_kv, seqused_cmp_kv, cmp_residual_kv, sinks, metadata, kv_quant_mode, rope_head_dim,
32 softmax_scale, cmp_ratio, ori_mask_mode, cmp_mask_mode, ori_win_left, ori_win_right, layout_q, layout_kv,32 softmax_scale, cmp_ratio, ori_mask_mode, cmp_mask_mode, ori_win_left, ori_win_right, layout_q, layout_kv,
33 topk_value_mode, return_softmax_lse):33 topk_value_mode, return_softmax_lse):
34- npu_result, _ = torch.ops.npu_ops_transformer.npu_mixed_quant_sparse_flash_mla(34+ npu_result, _ = torch.ops.cann_ops_transformer.mixed_quant_sparse_flash_mla(
35 q=q,35 q=q,
36 ori_kv=ori_kv,36 ori_kv=ori_kv,
37 cmp_kv=cmp_kv,37 cmp_kv=cmp_kv,
@@ -86,7 +86,7 @@ def test_qsmla_quant_process_graph(test_data, device_id=0):
86 npu_mode = torch.compile(npu_mode, fullgraph=True, backend=npu_backend, dynamic=False)86 npu_mode = torch.compile(npu_mode, fullgraph=True, backend=npu_backend, dynamic=False)
87 87 
88 print("test_data:", params)88 print("test_data:", params)
89- print("npu_mixed_quant_sparse_flash_mla_metadata...")89+ print("mixed_quant_sparse_flash_mla_metadata...")
90 layout_kv_qsas = metadata_input['layout_kv']90 layout_kv_qsas = metadata_input['layout_kv']
91 if layout_kv_qsas == "PA_BBND":91 if layout_kv_qsas == "PA_BBND":
92 layout_kv_qsas = "PA_ND"92 layout_kv_qsas = "PA_ND"
@@ -118,7 +118,7 @@ def test_qsmla_quant_process_graph(test_data, device_id=0):
118 torch.npu.synchronize()118 torch.npu.synchronize()
119 metadata.npu()119 metadata.npu()
120 120 
121- print("npu_mixed_quant_sparse_flash_mla...")121+ print("mixed_quant_sparse_flash_mla...")
122 npu_result = npu_mode(122 npu_result = npu_mode(
123 q=op_input['q'].npu() if op_input['q'] is not None else None,123 q=op_input['q'].npu() if op_input['q'] is not None else None,
124 ori_kv=op_input['ori_kv'].npu() if op_input['ori_kv'] is not None else None,124 ori_kv=op_input['ori_kv'].npu() if op_input['ori_kv'] is not None else None,
@@ -147,7 +147,7 @@ def test_qsmla_quant_process_graph(test_data, device_id=0):
147 return_softmax_lse=op_input.get('return_softmax_lse', False))147 return_softmax_lse=op_input.get('return_softmax_lse', False))
148 148 
149 torch.npu.synchronize()149 torch.npu.synchronize()
150- print("npu_mixed_quant_sparse_flash_mla...")150+ print("mixed_quant_sparse_flash_mla...")
151 npu_result = npu_mode(151 npu_result = npu_mode(
152 q=input['q'].npu() if input['q'] is not None else None,152 q=input['q'].npu() if input['q'] is not None else None,
153 ori_kv=input['ori_kv'].npu() if input['ori_kv'] is not None else None,153 ori_kv=input['ori_kv'].npu() if input['ori_kv'] is not None else None,
@@ -179,7 +179,7 @@ def test_qsmla_quant_process_ci(test_data, device_id=0):
179 torch_npu.npu.set_device(device_id)179 torch_npu.npu.set_device(device_id)
180 180 
181 print("test_data:", params)181 print("test_data:", params)
182- print("npu_mixed_quant_sparse_flash_mla_metadata...")182+ print("mixed_quant_sparse_flash_mla_metadata...")
183 layout_kv_qsas = metadata_input['layout_kv']183 layout_kv_qsas = metadata_input['layout_kv']
184 if layout_kv_qsas == "PA_BBND":184 if layout_kv_qsas == "PA_BBND":
185 layout_kv_qsas = "PA_ND"185 layout_kv_qsas = "PA_ND"
@@ -210,8 +210,8 @@ def test_qsmla_quant_process_ci(test_data, device_id=0):
210 torch.npu.synchronize()210 torch.npu.synchronize()
211 metadata.npu()211 metadata.npu()
212 212 
213- print("npu_mixed_quant_sparse_flash_mla...")213+ print("mixed_quant_sparse_flash_mla...")
214- npu_result, _ = torch.ops.npu_ops_transformer.npu_mixed_quant_sparse_flash_mla(214+ npu_result, _ = torch.ops.cann_ops_transformer.mixed_quant_sparse_flash_mla(
215 q=op_input['q'].npu() if op_input['q'] is not None else None,215 q=op_input['q'].npu() if op_input['q'] is not None else None,
216 ori_kv=op_input['ori_kv'].npu() if op_input['ori_kv'] is not None else None,216 ori_kv=op_input['ori_kv'].npu() if op_input['ori_kv'] is not None else None,
217 cmp_kv=op_input['cmp_kv'].npu() if op_input['cmp_kv'] is not None else None,217 cmp_kv=op_input['cmp_kv'].npu() if op_input['cmp_kv'] is not None else None,
@@ -19,7 +19,7 @@ import numpy as np
19import math19import math
20import ctypes20import ctypes
21import copy21import copy
22-import npu_ops_transformer22+import cann_ops_transformer
23 23 
24FP32_FRACTION_BITS = 23 # fp32尾数位数24FP32_FRACTION_BITS = 23 # fp32尾数位数
25 25 
@@ -840,7 +840,7 @@ def qliv2_output_single(params):
840 block_table = torch.from_numpy(block_table).to(dtype=torch.int32).npu()840 block_table = torch.from_numpy(block_table).to(dtype=torch.int32).npu()
841 max_seqlen_q_meta = actual_seq_lengths_query.max().item()841 max_seqlen_q_meta = actual_seq_lengths_query.max().item()
842 max_seqlen_k_meta = actual_seq_lengths_key.max().item()842 max_seqlen_k_meta = actual_seq_lengths_key.max().item()
843- metadata = torch.ops.npu_ops_transformer.npu_quant_lightning_indexer_v2_metadata(843+ metadata = torch.ops.cann_ops_transformer.quant_lightning_indexer_metadata(
844 cu_seqlens_q = cu_seqlens_query,844 cu_seqlens_q = cu_seqlens_query,
845 cu_seqlens_k = cu_seqlens_key,845 cu_seqlens_k = cu_seqlens_key,
846 seqused_q = actual_seq_lengths_query,846 seqused_q = actual_seq_lengths_query,
@@ -860,7 +860,7 @@ def qliv2_output_single(params):
860 cmp_ratio = cmp_ratio)860 cmp_ratio = cmp_ratio)
861 861 
862 metadata = metadata.npu()862 metadata = metadata.npu()
863- npu_result, _ = torch.ops.npu_ops_transformer.npu_quant_lightning_indexer_v2(query, key, weights,863+ npu_result, _ = torch.ops.cann_ops_transformer.quant_lightning_indexer_v2(query, key, weights,
864 query_dequant_scale,864 query_dequant_scale,
865 key_dequant_scale,865 key_dequant_scale,
866 cu_seqlens_q = cu_seqlens_query,866 cu_seqlens_q = cu_seqlens_query,
@@ -163,7 +163,7 @@
163 <tr>163 <tr>
164 <td>metadata</td>164 <td>metadata</td>
165 <td>可选输入</td>165 <td>可选输入</td>
166- <td>aicpu算子(npu_sparse_flash_mla_metadata)的分核结果。</td>166+ <td>aicpu算子(sparse_flash_mla_metadata)的分核结果。</td>
167 <td>INT32</td>167 <td>INT32</td>
168 <td>ND</td>168 <td>ND</td>
169 </tr>169 </tr>
@@ -219,7 +219,7 @@
219 <tr>219 <tr>
220 <td>layout_kv</td>220 <td>layout_kv</td>
221 <td>可选属性</td>221 <td>可选属性</td>
222- <td>用于标识输入`ori_kv`和`cmp_kv`的数据排布格式,支持输入"PA_BNBD"和"BSND"。</td>222+ <td>用于标识输入`ori_kv`和`cmp_kv`的数据排布格式,支持输入"PA_BBND"和"BSND"。</td>
223 <td>STRING</td>223 <td>STRING</td>
224 <td>-</td>224 <td>-</td>
225 </tr>225 </tr>
@@ -267,8 +267,8 @@
267 - `ori_kv``cmp_kv`的shape分别为[ori\_block\_num, ori\_block\_size, KV\_N, D]和[cmp\_block\_num, cmp\_block\_size, KV\_N, D],其中ori\_block\_num和cmp\_block\_num为PageAttention时block总数,ori\_block\_size和cmp\_block\_size为一个block的token数,ori\_block\_size和cmp\_block\_size取值为16的倍数,最大支持1024,KV_N仅支持1。267 - `ori_kv``cmp_kv`的shape分别为[ori\_block\_num, ori\_block\_size, KV\_N, D]和[cmp\_block\_num, cmp\_block\_size, KV\_N, D],其中ori\_block\_num和cmp\_block\_num为PageAttention时block总数,ori\_block\_size和cmp\_block\_size为一个block的token数,ori\_block\_size和cmp\_block\_size取值为16的倍数,最大支持1024,KV_N仅支持1。
268 - `ori_block_table``cmp_block_table`的shape为2维,其中第一维长度为B,第二维长度不小于所有batch中最大的S2和S3对应的block数量,即S2\_max / block\_size和S3\_max / block\_size向上取整。268 - `ori_block_table``cmp_block_table`的shape为2维,其中第一维长度为B,第二维长度不小于所有batch中最大的S2和S3对应的block数量,即S2\_max / block\_size和S3\_max / block\_size向上取整。
269- `metadata`为算子实际需要使用的分核结果,目前该参数必传,shape大小固定为[1024]。269- `metadata`为算子实际需要使用的分核结果,目前该参数必传,shape大小固定为[1024]。
270-- `layout_kv`仅支持输入"PA_BNBD"和"BSND"。270+- `layout_kv`仅支持输入"PA_BBND"和"BSND"。
271- - 当输入为PA_BNBD时,设置`cu_seqlens_ori_kv`和`cu_seqlens_cmp_kv`无效。271+ - 当输入为PA_BBND时,设置`cu_seqlens_ori_kv`和`cu_seqlens_cmp_kv`无效。
272 - 当输入为BSND时,`ori_kv``cmp_kv`的layout都必须为BSND,ori_kv的shape为[B, S2, N2,D],cmp_kv的shape为[B, S3, N2,D]。272 - 当输入为BSND时,`ori_kv``cmp_kv`的layout都必须为BSND,ori_kv的shape为[B, S2, N2,D],cmp_kv的shape为[B, S3, N2,D]。
273- 目前暂不支持返回`softmax_lse``return_softmax_lse`仅支持输入False,返回值`softmax_lse`为无效值。273- 目前暂不支持返回`softmax_lse``return_softmax_lse`仅支持输入False,返回值`softmax_lse`为无效值。
274- 目前暂不支持指定`q`中参与运算的token数,因此设置`seqused_q`无效。274- 目前暂不支持指定`q`中参与运算的token数,因此设置`seqused_q`无效。
@@ -19,13 +19,13 @@ import random
19import torch19import torch
20import torch_npu20import torch_npu
21 21 
22-# Register npu_sparse_flash_mla and npu_quant_lightning_indexer_v2_metadata via PTA22+# Register sparse_flash_mla and sparse_flash_mla_metadata via PTA
23TORCH_EXT_PATH = os.path.normpath(os.path.join(os.path.dirname(os.path.abspath(__file__)),23TORCH_EXT_PATH = os.path.normpath(os.path.join(os.path.dirname(os.path.abspath(__file__)),
24 '../../../../torch_extension'))24 '../../../../torch_extension'))
25if TORCH_EXT_PATH not in sys.path:25if TORCH_EXT_PATH not in sys.path:
26 sys.path.insert(0, TORCH_EXT_PATH)26 sys.path.insert(0, TORCH_EXT_PATH)
27-from npu_ops_transformer.ops import sparse_flash_mla as _smla_registration # noqa: F40127+from cann_ops_transformer.ops import sparse_flash_mla as _smla_registration # noqa: F401
28-from npu_ops_transformer.ops import quant_lightning_indexer_v2_metadata as _metadata_registration # noqa: F40128+from cann_ops_transformer.ops import sparse_flash_mla_metadata as _metadata_registration # noqa: F401
O
OopenLiBingCI6月16日

【openlibing.ci】识别到代码检查告警抑制注释,匹配工具:flake8,请Committer检视其合理性。

likedislike
29 29 
30class Network(torch.nn.Module):30class Network(torch.nn.Module):
O
OopenLiBingCI6月16日

【openlibing.ci】识别到代码检查告警抑制注释,匹配工具:flake8,请Committer检视其合理性。

likedislike
31 def __init__(self):31 def __init__(self):
@@ -89,8 +89,8 @@ def call_npu(input_data):
89 cmp_block_table = cmp_block_table.npu()89 cmp_block_table = cmp_block_table.npu()
90 90 
91 # 生成 metadata91 # 生成 metadata
92- print("npu_sparse_flash_mla_metadata...")92+ print("sparse_flash_mla_metadata...")
93- metadata = torch.ops.npu_ops_transformer.npu_sparse_flash_mla_metadata(93+ metadata = torch.ops.cann_ops_transformer.sparse_flash_mla_metadata(
94 num_heads_q=N1,94 num_heads_q=N1,
95 num_heads_kv=N2,95 num_heads_kv=N2,
96 head_dim=D,96 head_dim=D,
@@ -107,8 +107,8 @@ def call_npu(input_data):
107 max_seqlen_q=max_seqlen_q,107 max_seqlen_q=max_seqlen_q,
108 # max_seqlen_ori_kv=max_seqlen_ori_kv,108 # max_seqlen_ori_kv=max_seqlen_ori_kv,
109 # max_seqlen_cmp_kv=max_seqlen_cmp_kv,109 # max_seqlen_cmp_kv=max_seqlen_cmp_kv,
110- ori_topk=K,110+ # ori_topk=K,
111- # cmp_topk=cmp_topk,111+ cmp_topk=K,
112 cmp_ratio=cmp_ratio if cmp_ratio is not None else 1,112 cmp_ratio=cmp_ratio if cmp_ratio is not None else 1,
113 ori_mask_mode=ori_mask_mode,113 ori_mask_mode=ori_mask_mode,
114 cmp_mask_mode=cmp_mask_mode if cmp_mask_mode is not None else 3,114 cmp_mask_mode=cmp_mask_mode if cmp_mask_mode is not None else 3,
@@ -117,12 +117,11 @@ def call_npu(input_data):
117 layout_q=layout_q,117 layout_q=layout_q,
118 layout_kv=layout_kv,118 layout_kv=layout_kv,
119 has_ori_kv=ori_kv != None,119 has_ori_kv=ori_kv != None,
120- has_cmp_kv=cmp_kv != None,120+ has_cmp_kv=cmp_kv != None)
121- batch_invariant=0)
122 121 
123- # 统一调用 npu_sparse_flash_mla,所有参数可选带默认值122+ # 统一调用 sparse_flash_mla,所有参数可选带默认值
124- print("npu_sparse_flash_mla...")123+ print("sparse_flash_mla...")
125- npu_result, softmax_lse = torch.ops.npu_ops_transformer.npu_sparse_flash_mla(q,124+ npu_result, softmax_lse = torch.ops.cann_ops_transformer.sparse_flash_mla(q,
126 ori_kv=ori_kv,125 ori_kv=ori_kv,
127 cmp_kv=cmp_kv,126 cmp_kv=cmp_kv,
128 ori_sparse_indices=ori_sparse_indices,127 ori_sparse_indices=ori_sparse_indices,
@@ -149,9 +148,8 @@ def call_npu(input_data):
149 layout_q=layout_q,148 layout_q=layout_q,
150 layout_kv=layout_kv,149 layout_kv=layout_kv,
151 topk_value_mode=1,150 topk_value_mode=1,
152- return_softmax_lse=return_softmax_lse if return_softmax_lse is not None else False,151+ return_softmax_lse=return_softmax_lse if return_softmax_lse is not None else False)
153- batch_invariant=0)152+ print("sparse_flash_mla end")
154- print("npu_sparse_flash_mla end")
155 153 
156 torch.npu.synchronize()154 torch.npu.synchronize()
157 return npu_result, softmax_lse155 return npu_result, softmax_lse
@@ -58,8 +58,8 @@ TEST_PARAMS = {
58 "ori_win_left": [127],58 "ori_win_left": [127],
59 "ori_win_right": [0]59 "ori_win_right": [0]
60 },60 },
61- "cfa_prefill_tnd_padding": {61+ "cfa_prefill_tnd": {
62- "testcase_name": ["cfa_prefill_tnd_padding"],62+ "testcase_name": ["cfa_prefill_tnd"],
63 "layout_q": ["TND"],63 "layout_q": ["TND"],
64 "layout_kv": ["TND"],64 "layout_kv": ["TND"],
65 "q_type": [torch.bfloat16],65 "q_type": [torch.bfloat16],
@@ -69,13 +69,13 @@ TEST_PARAMS = {
69 "S1": [1],69 "S1": [1],
70 "T1": [1],70 "T1": [1],
71 "T2": [18],71 "T2": [18],
72- "T3": [5],72+ "T3": [4],
73 "N1": [64],73 "N1": [64],
74 "N2": [1],74 "N2": [1],
75 "D": [512],75 "D": [512],
76 "cu_seqlens_q": [[0, 1]],76 "cu_seqlens_q": [[0, 1]],
77 "cu_seqlens_ori_kv": [[0, 18]],77 "cu_seqlens_ori_kv": [[0, 18]],
78- "cu_seqlens_cmp_kv": [[0, 5]],78+ "cu_seqlens_cmp_kv": [[0, 4]],
79 "seqused_cmp_kv": [[4]],79 "seqused_cmp_kv": [[4]],
80 "cmp_residual_kv": [[2]],80 "cmp_residual_kv": [[2]],
81 "softmax_scale": [0.04419417],81 "softmax_scale": [0.04419417],
@@ -33,7 +33,7 @@ MINDSPEED_PATH = REPO_ROOT / "MindSpeed"
33if MINDSPEED_PATH.exists():33if MINDSPEED_PATH.exists():
34 sys.path.insert(0, str(MINDSPEED_PATH))34 sys.path.insert(0, str(MINDSPEED_PATH))
35 35 
36-import npu_ops_transformer # noqa: E402,F401 pylint: disable=wrong-import-position,unused-import36+import cann_ops_transformer # noqa: E402,F401 pylint: disable=wrong-import-position,unused-import
O
OopenLiBingCI6月16日

【openlibing.ci】识别到代码检查告警抑制注释,匹配工具:flake8,请Committer检视其合理性。

likedislike
37 37 
38 38 
39SLI_METADATA_SIZE = 6439SLI_METADATA_SIZE = 64
@@ -462,7 +462,7 @@ def run_fused_op(case, tensors, golden, device, cmp_ratio):
462 print("layout:", case.layout, "mask_mode:", case.sparse_mode, "cmp_ratio:", cmp_ratio)462 print("layout:", case.layout, "mask_mode:", case.sparse_mode, "cmp_ratio:", cmp_ratio)
463 print("##################### END fused inputs #######################")463 print("##################### END fused inputs #######################")
464 464 
465- metadata = torch.ops.npu_ops_transformer.npu_sparse_lightning_indexer_kl_loss_grad_metadata(465+ metadata = torch.ops.cann_ops_transformer.npu_sparse_lightning_indexer_kl_loss_grad_metadata(
466 case.n_query_index,466 case.n_query_index,
467 case.n_key,467 case.n_key,
468 case.d_index,468 case.d_index,
@@ -480,7 +480,7 @@ def run_fused_op(case, tensors, golden, device, cmp_ratio):
480 )480 )
481 print_slig_metadata(metadata)481 print_slig_metadata(metadata)
482 482 
483- dq, dk, dw, softmax_out = torch.ops.npu_ops_transformer.npu_sparse_lightning_indexer_kl_loss_grad(483+ dq, dk, dw, softmax_out = torch.ops.cann_ops_transformer.npu_sparse_lightning_indexer_kl_loss_grad(
484 q,484 q,
485 k,485 k,
486 w,486 w,
@@ -127,9 +127,9 @@ function(pack_built_in)
127 DESTINATION share/info/ops_transformer/script127 DESTINATION share/info/ops_transformer/script
128 )128 )
129 129 
130- # 打包 npu_ops_transformer whl 文件130+ # 打包 cann_ops_transformer whl 文件
131 set(WHL_SOURCE_DIR "${CMAKE_SOURCE_DIR}/torch_extension/dist")131 set(WHL_SOURCE_DIR "${CMAKE_SOURCE_DIR}/torch_extension/dist")
132- file(GLOB WHL_FILES "${WHL_SOURCE_DIR}/npu_ops_transformer-*.whl")132+ file(GLOB WHL_FILES "${WHL_SOURCE_DIR}/cann_ops_transformer-*.whl")
133 if(WHL_FILES)133 if(WHL_FILES)
134 install(FILES ${WHL_FILES}134 install(FILES ${WHL_FILES}
135 DESTINATION ${WHL_INSTALL_DIR}/es_packages/whl135 DESTINATION ${WHL_INSTALL_DIR}/es_packages/whl
@@ -18,4 +18,4 @@
18 18 
19| 接口名 | 说明 | 确定性说明 |19| 接口名 | 说明 | 确定性说明 |
20| ----------- | ------------------- | ------------------- |20| ----------- | ------------------- | ------------------- |
21-|[flash_attn](../../torch_extension/npu_ops_transformer/doc/npu_flash_attn.md)|完成xx计算。|xx|21+|[flash_attn](../../torch_extension/cann_ops_transformer/doc/npu_flash_attn.md)|完成xx计算。|xx|
@@ -284,4 +284,4 @@
284 284 
285| 调用方式 | 样例代码 | 说明 |285| 调用方式 | 样例代码 | 说明 |
286| :--------: | :----------------------------------------: | :-------------------------------------------------------: |286| :--------: | :----------------------------------------: | :-------------------------------------------------------: |
287-| PyTorch接口调用 | [deepep.py](../../torch_extension/npu_ops_transformer/ops/deep_ep.py) | 通过[mega_moe](../../torch_extension/npu_ops_transformer/doc/mega_moe.md)PyTorch接口方式调用mega_moe算子。 |287+| PyTorch接口调用 | [deepep.py](../../torch_extension/cann_ops_transformer/ops/deep_ep.py) | 通过[mega_moe](../../torch_extension/cann_ops_transformer/doc/mega_moe.md)PyTorch接口方式调用mega_moe算子。 |
@@ -368,10 +368,10 @@
368 - `shared_expert_num`:当前取值范围[0, 4]。368 - `shared_expert_num`:当前取值范围[0, 4]。
369 - `comm_quant_mode`:int8量化当且仅当`tp_world_size` < 2时可开启。369 - `comm_quant_mode`:int8量化当且仅当`tp_world_size` < 2时可开启。
370 - `performance_info_optional`:预留参数,当前版本不支持,传空指针即可。370 - `performance_info_optional`:预留参数,当前版本不支持,传空指针即可。
371- - `ccl_buffer_size`:调用get_low_latency_ccl_buffer_size接口(../../torch_extension/npu_ops_transformer/ops/deep_ep.py)。371+ - `ccl_buffer_size`:调用get_low_latency_ccl_buffer_size接口(../../torch_extension/cann_ops_transformer/ops/deep_ep.py)。
372 372 
373## 调用说明373## 调用说明
374 374 
375| 调用方式 | 样例代码 | 说明 |375| 调用方式 | 样例代码 | 说明 |
376| :--------: | :----------------------------------------: | :-------------------------------------------------------: |376| :--------: | :----------------------------------------: | :-------------------------------------------------------: |
377-| npu_low_latency_combine接口 | [deepep.py](../../torch_extension/npu_ops_transformer/ops/deep_ep.py) | 通过[npu_low_latency_combine](../../torch_extension/npu_ops_transformer/doc/npu_low_latency_combine.md)接口方式调用moe_distribute_combine_v3算子。 |377+| npu_low_latency_combine接口 | [deepep.py](../../torch_extension/cann_ops_transformer/ops/deep_ep.py) | 通过[npu_low_latency_combine](../../torch_extension/cann_ops_transformer/doc/npu_low_latency_combine.md)接口方式调用moe_distribute_combine_v3算子。 |
@@ -343,7 +343,7 @@ $$
343 - "fullmesh_v2":开启fullmesh_v2模板,其中`comm_alg`仅在`tp_world_size`取值为1时生效,且不支持在各卡`BS`不一致、输入xActiveMask和特殊专家场景下开启。343 - "fullmesh_v2":开启fullmesh_v2模板,其中`comm_alg`仅在`tp_world_size`取值为1时生效,且不支持在各卡`BS`不一致、输入xActiveMask和特殊专家场景下开启。
344 - `ep_recv_count_out`:要求shape为(`ep_world_size` * max(`tp_world_size`, 1) * `local_expert_num`, )。344 - `ep_recv_count_out`:要求shape为(`ep_world_size` * max(`tp_world_size`, 1) * `local_expert_num`, )。
345 - `performance_Info_optional`:预留参数,当前版本不支持,传空指针即可。345 - `performance_Info_optional`:预留参数,当前版本不支持,传空指针即可。
346- - `ccl_buffer_size`:调用get_low_latency_ccl_buffer_size接口(../../torch_extension/npu_ops_transformer/ops/deep_ep.py)。346+ - `ccl_buffer_size`:调用get_low_latency_ccl_buffer_size接口(../../torch_extension/cann_ops_transformer/ops/deep_ep.py)。
347 - 参数说明里shape格式说明:347 - 参数说明里shape格式说明:
348 - `H`:表示hidden size隐藏层大小,取值范围[1024, 8192]。348 - `H`:表示hidden size隐藏层大小,取值范围[1024, 8192]。
349 - `BS`:表示batch sequence size,即本卡最终输出的token数量,取值范围为[1, 512]。349 - `BS`:表示batch sequence size,即本卡最终输出的token数量,取值范围为[1, 512]。
@@ -352,4 +352,4 @@ $$
352 352 
353| 调用方式 | 样例代码 | 说明 |353| 调用方式 | 样例代码 | 说明 |
354| :--------: | :----------------------------------------: | :-------------------------------------------------------: |354| :--------: | :----------------------------------------: | :-------------------------------------------------------: |
355-| npu_low_latency_dispatch接口 | [deepep.py](../../torch_extension/npu_ops_transformer/ops/deep_ep.py) | 通过[npu_low_latency_dispatch](../../torch_extension/npu_ops_transformer/doc/npu_low_latency_dispatch.md)接口方式调用moe_distribute_dispatch_v3算子。 |355+| npu_low_latency_dispatch接口 | [deepep.py](../../torch_extension/cann_ops_transformer/ops/deep_ep.py) | 通过[npu_low_latency_dispatch](../../torch_extension/cann_ops_transformer/doc/npu_low_latency_dispatch.md)接口方式调用moe_distribute_dispatch_v3算子。 |
@@ -375,11 +375,11 @@ install_whl_package() {
375 local target_python_dir="${TARGET_VERSION_DIR}/python/site-packages"375 local target_python_dir="${TARGET_VERSION_DIR}/python/site-packages"
376 376 
377 # 查找 whl 文件377 # 查找 whl 文件
378- local whl_file=$(find "${whl_dir}" -name "npu_ops_transformer-*.whl" 2>/dev/null | head -1)378+ local whl_file=$(find "${whl_dir}" -name "cann_ops_transformer-*.whl" 2>/dev/null | head -1)
379 379 
380 if [ -n "${whl_file}" ] && [ -f "${whl_file}" ]; then380 if [ -n "${whl_file}" ] && [ -f "${whl_file}" ]; then
381 logandprint "[INFO]: Found whl package: ${whl_file}"381 logandprint "[INFO]: Found whl package: ${whl_file}"
382- logandprint "[INFO]: Installing npu_ops_transformer whl package to ${target_python_dir}"382+ logandprint "[INFO]: Installing cann_ops_transformer whl package to ${target_python_dir}"
383 383 
384 # 创建目标目录384 # 创建目标目录
385 comm_create_dir "${target_python_dir}" "${CREATE_DIR_PERM}" "${TARGET_USERNAME}:${TARGET_USERGROUP}" "${IS_FOR_ALL}"385 comm_create_dir "${target_python_dir}" "${CREATE_DIR_PERM}" "${TARGET_USERNAME}:${TARGET_USERGROUP}" "${IS_FOR_ALL}"
@@ -403,9 +403,9 @@ install_whl_package() {
403 cp "${whl_file}" "${target_python_dir}/"403 cp "${whl_file}" "${target_python_dir}/"
404 fi404 fi
405 405 
406- logandprint "[INFO]: npu_ops_transformer whl package installed successfully"406+ logandprint "[INFO]: cann_ops_transformer whl package installed successfully"
407 else407 else
408- logandprint "[INFO]: No npu_ops_transformer whl package found, skipping"408+ logandprint "[INFO]: No cann_ops_transformer whl package found, skipping"
409 fi409 fi
410}410}
411 411 
@@ -183,19 +183,19 @@ remove_init_py() {
183remove_whl_package() {183remove_whl_package() {
184 local python_dir="${TARGET_VERSION_DIR}/python/site-packages"184 local python_dir="${TARGET_VERSION_DIR}/python/site-packages"
185 185 
186- if [ -d "${python_dir}/npu_ops_transformer" ]; then186+ if [ -d "${python_dir}/cann_ops_transformer" ]; then
187- logandprint "[INFO]: Removing npu_ops_transformer whl package from ${python_dir}"187+ logandprint "[INFO]: Removing cann_ops_transformer whl package from ${python_dir}"
188 188 
189 # 尝试使用 pip 卸载189 # 尝试使用 pip 卸载
190 if command -v pip3 &>/dev/null; then190 if command -v pip3 &>/dev/null; then
191- pip3 uninstall -y npu_ops_transformer --target="${python_dir}" 2>/dev/null191+ pip3 uninstall -y cann_ops_transformer --target="${python_dir}" 2>/dev/null
192 fi192 fi
193 193 
194 # 直接删除目录(确保清理干净)194 # 直接删除目录(确保清理干净)
195- rm -rf "${python_dir}/npu_ops_transformer" 2>/dev/null195+ rm -rf "${python_dir}/cann_ops_transformer" 2>/dev/null
196- rm -f ${python_dir}/npu_ops_transformer-*.whl 2>/dev/null196+ rm -f ${python_dir}/cann_ops_transformer-*.whl 2>/dev/null
197- rm -rf ${python_dir}/npu_ops_transformer-*.egg-info 2>/dev/null197+ rm -rf ${python_dir}/cann_ops_transformer-*.egg-info 2>/dev/null
198- rm -rf ${python_dir}/npu_ops_transformer-*.dist-info 2>/dev/null198+ rm -rf ${python_dir}/cann_ops_transformer-*.dist-info 2>/dev/null
199 199 
200 # 如果目录为空,删除它200 # 如果目录为空,删除它
201 if [ -d "${python_dir}" ] && [ -z "$(ls -A ${python_dir} 2>/dev/null)" ]; then201 if [ -d "${python_dir}" ] && [ -z "$(ls -A ${python_dir} 2>/dev/null)" ]; then
@@ -207,7 +207,7 @@ remove_whl_package() {
207 fi207 fi
208 fi208 fi
209 209 
210- logandprint "[INFO]: npu_ops_transformer whl package removed"210+ logandprint "[INFO]: cann_ops_transformer whl package removed"
211 fi211 fi
212}212}
213 213 
@@ -1,6 +1,6 @@
1-# NPU Ops Transformer1+# CANN Ops Transformer
2 2 
3-`npu_ops_transformer` is a high-performance operator extension library designed for Ascend NPU. It leverages Just-In-Time(JIT) compilation to bridge PyTorch functional interfaces with ACLNN library.3+`cann_ops_transformer` is a high-performance operator extension library designed for Ascend NPU. It leverages Just-In-Time(JIT) compilation to bridge PyTorch functional interfaces with ACLNN library.
4 4 
5## Build & Installation5## Build & Installation
6 6 
@@ -37,19 +37,19 @@
37 37 
38## Quick Start38## Quick Start
39 39 
40-Using `npu_ops_transformer` is seamless. You can invoke NPU-accelerated operators directly through the library's opset.40+Using `cann_ops_transformer` is seamless. You can invoke NPU-accelerated operators directly through the library's opset.
41 41 
42```python42```python
43import torch43import torch
44import torch_npu44import torch_npu
45-import npu_ops_transformer45+import cann_ops_transformer
46 46 
47# Initialize data on NPU47# Initialize data on NPU
48x = torch.randn(10, 32, dtype=torch.float32).npu()48x = torch.randn(10, 32, dtype=torch.float32).npu()
49 49 
50# Call the custom NPU operator50# Call the custom NPU operator
51# This triggers JIT compilation on the first call51# This triggers JIT compilation on the first call
52-npu_result = npu_ops_transformer.ops.abs(x)52+npu_result = cann_ops_transformer.ops.abs(x)
53 53 
54# Verify against CPU ATen implementation54# Verify against CPU ATen implementation
55cpu_x = x.cpu()55cpu_x = x.cpu()
@@ -102,8 +102,8 @@ This file manages the JIT compilation logic and registers the operator into the
102import torch102import torch
103import torch_npu103import torch_npu
104from torch.library import impl104from torch.library import impl
105-from npu_ops_transformer.op_builder.builder import OpBuilder105+from cann_ops_transformer.op_builder.builder import OpBuilder
106-from npu_ops_transformer.op_builder.builder import AS_LIBRARY106+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
107 107 
108class AbsOpBuilder(OpBuilder):108class AbsOpBuilder(OpBuilder):
109 def __init__(self):109 def __init__(self):
Rtorch_extension/npu_ops_transformer/__init__.pytorch_extension/cann_ops_transformer/__init__.py+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/common/inc/aclnn_common.htorch_extension/cann_ops_transformer/common/inc/aclnn_common.h+10-5
@@ -6,15 +6,15 @@
6 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,6 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.7 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8 * See LICENSE in the root of the software repository for the full text of the License.8 * See LICENSE in the root of the software repository for the full text of the License.
9- */9+ */
10 10 
11/*!11/*!
12 * \file aclnn_common.h12 * \file aclnn_common.h
13 * \brief13 * \brief
14 */14 */
15 15 
16-#ifndef NPU_OPS_TRANSFORMER_ACLNN_COMMON_H16+#ifndef CANN_OPS_TRANSFORMER_ACLNN_COMMON_H
17-#define NPU_OPS_TRANSFORMER_ACLNN_COMMON_H17+#define CANN_OPS_TRANSFORMER_ACLNN_COMMON_H
18 18 
19#include <torch_npu/csrc/framework/utils/OpAdapter.h>19#include <torch_npu/csrc/framework/utils/OpAdapter.h>
20#include <dlfcn.h>20#include <dlfcn.h>
@@ -773,6 +773,11 @@ inline void Release(aclTensorList *p)
773 aclDestroyTensorList(p);773 aclDestroyTensorList(p);
774}774}
775 775 
776+inline const c10::optional<at::Tensor> get_valid_tensor(const c10::optional<at::Tensor> &tensor_opt, at::Device device)
777+{
778+ return tensor_opt.has_value() ? tensor_opt : torch::empty({0}, torch::dtype(torch::kInt32).device(device));
779+};
780+ 
776template <typename T>781template <typename T>
777void Release(T value)782void Release(T value)
778{783{
@@ -935,7 +940,7 @@ auto DecodeDevice(Ts&... args) -> at::Device
935 workspace_tensor = at::empty({workspace_size}, options.dtype(at::kByte)); \940 workspace_tensor = at::empty({workspace_size}, options.dtype(at::kByte)); \
936 workspace_addr = const_cast<void *>(workspace_tensor.storage().data()); \941 workspace_addr = const_cast<void *>(workspace_tensor.storage().data()); \
937 } \942 } \
938- auto acl_call = [converted_params, workspace_addr, workspace_size, acl_stream, executor]() -> int { \943+ auto acl_call = [converted_params, workspace_addr, workspace_size, acl_stream, executor]()->int { \
939 typedef int (*OpApiFunc)(void *, uint64_t, aclOpExecutor *, const aclrtStream); \944 typedef int (*OpApiFunc)(void *, uint64_t, aclOpExecutor *, const aclrtStream); \
940 OpApiFunc opApiFunc = reinterpret_cast<OpApiFunc>(opApiFuncAddr); \945 OpApiFunc opApiFunc = reinterpret_cast<OpApiFunc>(opApiFuncAddr); \
941 auto api_ret = opApiFunc(workspace_addr, workspace_size, executor, acl_stream); \946 auto api_ret = opApiFunc(workspace_addr, workspace_size, executor, acl_stream); \
@@ -956,4 +961,4 @@ auto DecodeDevice(Ts&... args) -> at::Device
956 } \ 961 } \
957 } while (false)962 } while (false)
958 963 
959-#endif // NPU_OPS_TRANSFORMER_ACLNN_COMMON_H964+#endif // CANN_OPS_TRANSFORMER_ACLNN_COMMON_H
Rtorch_extension/npu_ops_transformer/common/inc/hccl_common.htorch_extension/cann_ops_transformer/common/inc/hccl_common.h+7-5
@@ -13,8 +13,8 @@
13 * \brief13 * \brief
14 */14 */
15 15 
16-#ifndef NPU_OPS_TRANSFORMER_HCCL_COMMON_H16+#ifndef CANN_OPS_TRANSFORMER_HCCL_COMMON_H
17-#define NPU_OPS_TRANSFORMER_HCCL_COMMON_H17+#define CANN_OPS_TRANSFORMER_HCCL_COMMON_H
18 18 
19#include "aclnn_common.h"19#include "aclnn_common.h"
20#include <torch_npu/csrc/framework/utils/OpAdapter.h>20#include <torch_npu/csrc/framework/utils/OpAdapter.h>
@@ -142,7 +142,8 @@ inline void InitHcclFunctions()
142 TORCH_CHECK(HcclKfcAllocOpArgsFunc != nullptr, "getHcclKfcAllocOpArgs failed.");142 TORCH_CHECK(HcclKfcAllocOpArgsFunc != nullptr, "getHcclKfcAllocOpArgs failed.");
143 HcclKfcFreeOpArgsFunc = GetHcclFuncAddr<_HcclKfcFreeOpArgs>("HcclKfcFreeOpArgs"); // 释放通信配置对象143 HcclKfcFreeOpArgsFunc = GetHcclFuncAddr<_HcclKfcFreeOpArgs>("HcclKfcFreeOpArgs"); // 释放通信配置对象
144 TORCH_CHECK(HcclKfcFreeOpArgsFunc != nullptr, "getHcclKfcFreeOpArgs failed.");144 TORCH_CHECK(HcclKfcFreeOpArgsFunc != nullptr, "getHcclKfcFreeOpArgs failed.");
145- HcclKfcOpArgsSetCommEngineFunc = GetHcclFuncAddr<_HcclKfcOpArgsSetCommEngine>("HcclKfcOpArgsSetCommEngine"); // 设置通信方式145+ HcclKfcOpArgsSetCommEngineFunc =
146+ GetHcclFuncAddr<_HcclKfcOpArgsSetCommEngine>("HcclKfcOpArgsSetCommEngine"); // 设置通信方式
146 TORCH_CHECK(HcclKfcOpArgsSetCommEngineFunc != nullptr, "getHcclKfcOpArgsSetCommEngine failed.");147 TORCH_CHECK(HcclKfcOpArgsSetCommEngineFunc != nullptr, "getHcclKfcOpArgsSetCommEngine failed.");
147 HcclGetRankIdFunc = GetHcclFuncAddr<_HcclGetRankId>("HcclGetRankId"); // 获取本卡卡号148 HcclGetRankIdFunc = GetHcclFuncAddr<_HcclGetRankId>("HcclGetRankId"); // 获取本卡卡号
148 TORCH_CHECK(HcclGetRankIdFunc != nullptr, "getFuncHcclGetRankId failed.");149 TORCH_CHECK(HcclGetRankIdFunc != nullptr, "getFuncHcclGetRankId failed.");
@@ -152,7 +153,8 @@ inline void InitHcclFunctions()
152 TORCH_CHECK(HcclGetRemoteIpcHcclBufFunc != nullptr, "getFuncHcclGetRemoteIpcHcclBuf failed.");153 TORCH_CHECK(HcclGetRemoteIpcHcclBufFunc != nullptr, "getFuncHcclGetRemoteIpcHcclBuf failed.");
153 HcclKfcOpArgsSetAlgConfigFunc = GetHcclFuncAddr<_HcclKfcOpArgsSetAlgConfig>("HcclKfcOpArgsSetAlgConfig"); // 设置通信类型154 HcclKfcOpArgsSetAlgConfigFunc = GetHcclFuncAddr<_HcclKfcOpArgsSetAlgConfig>("HcclKfcOpArgsSetAlgConfig"); // 设置通信类型
154 TORCH_CHECK(HcclKfcOpArgsSetAlgConfigFunc != nullptr, "getFuncHcclKfcOpArgsSetAlgConfig failed.");155 TORCH_CHECK(HcclKfcOpArgsSetAlgConfigFunc != nullptr, "getFuncHcclKfcOpArgsSetAlgConfig failed.");
155- HcclCommGetHandleWithNameFunc = GetHcclFwkFuncAddr<_HcclCommGetHandleWithName>("HcclCommGetHandleWithName"); // 通过groupName获取groupHandle156+ HcclCommGetHandleWithNameFunc =
157+ GetHcclFwkFuncAddr<_HcclCommGetHandleWithName>("HcclCommGetHandleWithName"); // 通过groupName获取groupHandle
156 TORCH_CHECK(HcclCommGetHandleWithNameFunc != nullptr, "getFuncHcclCommGetHandleWithName failed.");158 TORCH_CHECK(HcclCommGetHandleWithNameFunc != nullptr, "getFuncHcclCommGetHandleWithName failed.");
157 HcclCreateOpResCtxFunc = GetHcclFuncAddr<_HcclCreateOpResCtx>("HcclCreateOpResCtx"); // 创建HcclContext159 HcclCreateOpResCtxFunc = GetHcclFuncAddr<_HcclCreateOpResCtx>("HcclCreateOpResCtx"); // 创建HcclContext
158 TORCH_CHECK(HcclCreateOpResCtxFunc != nullptr, "getFuncHcclCreateOpResCtx failed.");160 TORCH_CHECK(HcclCreateOpResCtxFunc != nullptr, "getFuncHcclCreateOpResCtx failed.");
@@ -193,4 +195,4 @@ inline void InitHcclEngineCtxFunctions()
193 TORCH_CHECK(HcclGetRankSizeFunc != nullptr, "getFuncHcclGetRankSize failed.");195 TORCH_CHECK(HcclGetRankSizeFunc != nullptr, "getFuncHcclGetRankSize failed.");
194}196}
195 197 
196-#endif // NPU_OPS_TRANSFORMER_HCCL_COMMON_H198+#endif // CANN_OPS_TRANSFORMER_HCCL_COMMON_H
Rtorch_extension/npu_ops_transformer/doc/get_low_latency_ccl_buffer_size.mdtorch_extension/cann_ops_transformer/doc/get_low_latency_ccl_buffer_size.md+1-1
@@ -76,7 +76,7 @@ get_low_latency_ccl_buffer_size(world_size, num_max_dispatch_tokens_per_rank, hi
76import os76import os
77import torch77import torch
78import torch_npu78import torch_npu
79-from npu_ops_transformer.ops import MoeDistributeBuffer79+from cann_ops_transformer.ops import MoeDistributeBuffer
80 80 
81server_num = 181server_num = 1
82dev_num = 1682dev_num = 16
Rtorch_extension/npu_ops_transformer/doc/mega_moe.mdtorch_extension/cann_ops_transformer/doc/mega_moe.md+1-1
@@ -181,7 +181,7 @@ get_symm_buffer_for_mega_moe(group, num_experts: int, num_max_tokens_per_rank: i
181 import torch.distributed as dist181 import torch.distributed as dist
182 from torch.distributed import ReduceOp182 from torch.distributed import ReduceOp
183 import torch.multiprocessing as mp183 import torch.multiprocessing as mp
184- from npu_ops_transformer.ops import get_symm_buffer_for_mega_moe, mega_moe184+ from cann_ops_transformer.ops import get_symm_buffer_for_mega_moe, mega_moe
185 import torchair185 import torchair
186 186 
187 E = 4187 E = 4
Rtorch_extension/npu_ops_transformer/doc/mhc_post.mdtorch_extension/cann_ops_transformer/doc/mhc_post.md+1-1
@@ -197,7 +197,7 @@ cann_ops_transformer.mhc_post(x, h_res, h_out, h_post) -> Tensor
197 ```python197 ```python
198 import torch198 import torch
199 import torch_npu199 import torch_npu
200- from npu_ops_transformer.ops import mhc_post200+ from cann_ops_transformer.ops import mhc_post
201 201 
202 B = 2202 B = 2
203 S = 8203 S = 8
Rtorch_extension/npu_ops_transformer/doc/mhc_pre_sinkhorn.mdtorch_extension/cann_ops_transformer/doc/mhc_pre_sinkhorn.md+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/doc/npu_flash_attn.mdtorch_extension/cann_ops_transformer/doc/npu_flash_attn.md+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/doc/npu_get_mega_moe_ccl_buffer_size.mdtorch_extension/cann_ops_transformer/doc/npu_get_mega_moe_ccl_buffer_size.md+1-1
@@ -68,7 +68,7 @@ npu_get_mega_moe_ccl_buffer_size(ep_world_size: int, moe_expert_num: int, num_ma
68import os68import os
69import torch69import torch
70import torch_npu70import torch_npu
71-from npu_ops_transformer.ops import npu_get_mega_moe_ccl_buffer_size71+from cann_ops_transformer.ops import npu_get_mega_moe_ccl_buffer_size
72 72 
73server_num = 173server_num = 1
74rank_per_dev = 274rank_per_dev = 2
Rtorch_extension/npu_ops_transformer/doc/npu_low_latency_combine.mdtorch_extension/cann_ops_transformer/doc/npu_low_latency_combine.md+2-2
@@ -148,7 +148,7 @@ npu_low_latency_combine(x, topk_idx, topk_weights, assist_info_for_combine, ep_s
148 from torch.multiprocessing import Process148 from torch.multiprocessing import Process
149 import torch.distributed as dist149 import torch.distributed as dist
150 from torch.distributed import ReduceOp150 from torch.distributed import ReduceOp
151- from npu_ops_transformer.ops import MoeDistributeBuffer151+ from cann_ops_transformer.ops import MoeDistributeBuffer
152 152 
153 # 控制模式153 # 控制模式
154 quant_mode = 2 # 2为动态量化154 quant_mode = 2 # 2为动态量化
@@ -352,7 +352,7 @@ npu_low_latency_combine(x, topk_idx, topk_weights, assist_info_for_combine, ep_s
352 import torch.distributed as dist352 import torch.distributed as dist
353 from torch.distributed import ReduceOp353 from torch.distributed import ReduceOp
354 import time354 import time
355- from npu_ops_transformer.ops import MoeDistributeBuffer355+ from cann_ops_transformer.ops import MoeDistributeBuffer
356 356 
357 # 控制模式357 # 控制模式
358 quant_mode = 2 # 2为动态量化358 quant_mode = 2 # 2为动态量化
Rtorch_extension/npu_ops_transformer/doc/npu_low_latency_dispatch.mdtorch_extension/cann_ops_transformer/doc/npu_low_latency_dispatch.md+2-2
@@ -165,7 +165,7 @@ npu_low_latency_dispatch(x, topk_idx, num_experts, *, quant_mode = 0, comm_alg="
165 from torch.multiprocessing import Process165 from torch.multiprocessing import Process
166 import torch.distributed as dist166 import torch.distributed as dist
167 from torch.distributed import ReduceOp167 from torch.distributed import ReduceOp
168- from npu_ops_transformer.ops import MoeDistributeBuffer168+ from cann_ops_transformer.ops import MoeDistributeBuffer
169 169 
170 # 控制模式170 # 控制模式
171 quant_mode = 2 # 2为动态量化171 quant_mode = 2 # 2为动态量化
@@ -369,7 +369,7 @@ npu_low_latency_dispatch(x, topk_idx, num_experts, *, quant_mode = 0, comm_alg="
369 import torch.distributed as dist369 import torch.distributed as dist
370 from torch.distributed import ReduceOp370 from torch.distributed import ReduceOp
371 import time371 import time
372- from npu_ops_transformer.ops import MoeDistributeBuffer372+ from cann_ops_transformer.ops import MoeDistributeBuffer
373 373 
374 # 控制模式374 # 控制模式
375 quant_mode = 2 # 2为动态量化375 quant_mode = 2 # 2为动态量化
Rtorch_extension/npu_ops_transformer/op_builder/builder.pytorch_extension/cann_ops_transformer/op_builder/builder.py+3-3
@@ -15,10 +15,10 @@ import torch
15from torch.utils.cpp_extension import load15from torch.utils.cpp_extension import load
16from torch.library import Library16from torch.library import Library
17import torch_npu17import torch_npu
18-import npu_ops_transformer18+import cann_ops_transformer
19 19 
20ASCEND_HOME_PATH = "ASCEND_HOME_PATH"20ASCEND_HOME_PATH = "ASCEND_HOME_PATH"
21-AS_LIBRARY = Library("npu_ops_transformer", "DEF")21+AS_LIBRARY = Library("cann_ops_transformer", "DEF")
22 22 
23 23 
24class OpBuilder(ABC):24class OpBuilder(ABC):
@@ -34,7 +34,7 @@ class OpBuilder(ABC):
34 self.name = name34 self.name = name
35 self._cann_path = self.get_cann_path()35 self._cann_path = self.get_cann_path()
36 self._torch_npu_path = os.path.dirname(os.path.abspath(torch_npu.__file__))36 self._torch_npu_path = os.path.dirname(os.path.abspath(torch_npu.__file__))
37- self._package_path = os.path.dirname(os.path.abspath(npu_ops_transformer.__file__))37+ self._package_path = os.path.dirname(os.path.abspath(cann_ops_transformer.__file__))
38 self.register_schema(self.schema())38 self.register_schema(self.schema())
39 self.register_meta()39 self.register_meta()
40 40 
Rtorch_extension/npu_ops_transformer/ops/__init__.pytorch_extension/cann_ops_transformer/ops/__init__.py+12-9
@@ -1,5 +1,5 @@
1# -----------------------------------------------------------------------------------------------------------1# -----------------------------------------------------------------------------------------------------------
2-# Copyright (c) 2025 Huawei Technologies Co., Ltd.2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3# This program is free software, you can redistribute it and/or modify it under the terms and conditions of3# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4# CANN Open Software License Agreement Version 2.0 (the "License").4# CANN Open Software License Agreement Version 2.0 (the "License").
5# Please refer to the License for details. You may not use this file except in compliance with the License.5# Please refer to the License for details. You may not use this file except in compliance with the License.
@@ -23,16 +23,19 @@ from .graph_convert.graph_convert_mega_moe import convert_npu_mega_moe
23from .flash_attn import npu_flash_attn23from .flash_attn import npu_flash_attn
24from .graph_convert.graph_convert_flash_attn import convert_npu_flash_attn24from .graph_convert.graph_convert_flash_attn import convert_npu_flash_attn
25from .flash_attn_metadata import npu_flash_attn_metadata25from .flash_attn_metadata import npu_flash_attn_metadata
26-from .lightning_indexer_v2_metadata import npu_lightning_indexer_v2_metadata26+from .mixed_quant_sparse_flash_mla import mixed_quant_sparse_flash_mla, mixed_quant_sparse_flash_mla_metadata
27-from .quant_lightning_indexer_v2_metadata import npu_quant_lightning_indexer_v2_metadata27+from .graph_convert.graph_convert_mixed_quant_sparse_flash_mla import (
28-from .mixed_quant_sparse_flash_mla import npu_mixed_quant_sparse_flash_mla28+ convert_mixed_quant_sparse_flash_mla_metadata
29-from .graph_convert.graph_convert_lightning_indexer_v2_metadata import convert_npu_lightning_indexer_v2_metadata
30-from .graph_convert.graph_convert_quant_lightning_indexer_v2_metadata import (
31- convert_npu_quant_lightning_indexer_v2_metadata
32)29)
30+from .sparse_flash_mla import sparse_flash_mla, sparse_flash_mla_metadata
31+from .graph_convert.graph_convert_sparse_flash_mla import convert_sparse_flash_mla_metadata
33from .npu_sparse_lightning_indexer_kl_loss_grad import npu_sparse_lightning_indexer_kl_loss_grad32from .npu_sparse_lightning_indexer_kl_loss_grad import npu_sparse_lightning_indexer_kl_loss_grad
34from .npu_sparse_lightning_indexer_kl_loss_grad_metadata import npu_sparse_lightning_indexer_kl_loss_grad_metadata33from .npu_sparse_lightning_indexer_kl_loss_grad_metadata import npu_sparse_lightning_indexer_kl_loss_grad_metadata
35-from .lightning_indexer_v2 import npu_lightning_indexer_v234+from .lightning_indexer_v2 import lightning_indexer_v2, lightning_indexer_metadata
36-from .quant_lightning_indexer_v2 import npu_quant_lightning_indexer_v235+from .graph_convert.graph_convert_lightning_indexer import convert_lightning_indexer_metadata
36+from .quant_lightning_indexer_v2 import quant_lightning_indexer_v2, quant_lightning_indexer_metadata
37+from .graph_convert.graph_convert_quant_lightning_indexer import (
38+ convert_quant_lightning_indexer_metadata
39+)
37from .mhc_post import mhc_post40from .mhc_post import mhc_post
38from .mhc_pre_sinkhorn import mhc_pre_sinkhorn41from .mhc_pre_sinkhorn import mhc_pre_sinkhorn
Rtorch_extension/npu_ops_transformer/ops/comm_context.pytorch_extension/cann_ops_transformer/ops/comm_context.py+1-1
@@ -9,7 +9,7 @@
9# -----------------------------------------------------------------------------------------------------------9# -----------------------------------------------------------------------------------------------------------
10import torch10import torch
11import torch_npu11import torch_npu
12-from npu_ops_transformer.op_builder.builder import OpBuilder12+from cann_ops_transformer.op_builder.builder import OpBuilder
13 13 
14 14 
15class CommContextOpBuilder(OpBuilder):15class CommContextOpBuilder(OpBuilder):
Rtorch_extension/npu_ops_transformer/ops/csrc/comm_context.cpptorch_extension/cann_ops_transformer/ops/csrc/comm_context.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/flash_attn.cpptorch_extension/cann_ops_transformer/ops/csrc/flash_attn.cpp+89-89
@@ -1,90 +1,90 @@
1-/**1+/**
2- * Copyright (c) 2026 Huawei Technologies Co., Ltd.2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3- * This program is free software, you can redistribute it and/or modify it under the terms and conditions of3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4- * CANN Open Software License Agreement Version 2.0 (the "License").4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5- * Please refer to the License for details. You may not use this file except in compliance with the License.5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6- * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7- * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8- * See LICENSE in the root of the software repository for the full text of the License.8+ * See LICENSE in the root of the software repository for the full text of the License.
9- */9+ */
10- 10+ 
11-/*!11+/*!
12- * \file npu_flash_attn.cpp12+ * \file npu_flash_attn.cpp
13- * \brief13+ * \brief
14- */14+ */
15- 15+ 
16-#include <torch/extension.h>16+#include <torch/extension.h>
17-#include "aclnn_common.h"17+#include "aclnn_common.h"
18- 18+ 
19-namespace op_api {19+namespace op_api {
20-const int64_t DIM_ONE = 1;20+const int64_t DIM_ONE = 1;
21-const int64_t DIM_TWO = 2;21+const int64_t DIM_TWO = 2;
22-const int64_t DIM_THREE = 3;22+const int64_t DIM_THREE = 3;
23-const int64_t MAX_DIM_SIZE = 8;23+const int64_t MAX_DIM_SIZE = 8;
24-std::tuple<at::Tensor, at::Tensor>24+std::tuple<at::Tensor, at::Tensor>
25-npu_flash_attn(const at::Tensor &q, const at::Tensor &k, const at::Tensor &v,25+npu_flash_attn(const at::Tensor &q, const at::Tensor &k, const at::Tensor &v,
26- const c10::optional<at::Tensor> &block_table, const c10::optional<at::Tensor> &cu_seqlens_q,26+ const c10::optional<at::Tensor> &block_table, const c10::optional<at::Tensor> &cu_seqlens_q,
27- const c10::optional<at::Tensor> &cu_seqlens_kv, const c10::optional<at::Tensor> &seqused_q,27+ const c10::optional<at::Tensor> &cu_seqlens_kv, const c10::optional<at::Tensor> &seqused_q,
28- const c10::optional<at::Tensor> &seqused_kv, const c10::optional<at::Tensor> &sinks,28+ const c10::optional<at::Tensor> &seqused_kv, const c10::optional<at::Tensor> &sinks,
29- const c10::optional<at::Tensor> &attn_mask, const c10::optional<at::Tensor> &metadata,29+ const c10::optional<at::Tensor> &attn_mask, const c10::optional<at::Tensor> &metadata,
30- double softmax_scale, int64_t mask_mode, int64_t win_left, int64_t win_right, int64_t max_seqlen_q,30+ double softmax_scale, int64_t mask_mode, int64_t win_left, int64_t win_right, int64_t max_seqlen_q,
31- int64_t max_seqlen_kv, string layout_q, string layout_kv, string layout_out, int64_t return_softmax_lse)31+ int64_t max_seqlen_kv, string layout_q, string layout_kv, string layout_out, int64_t return_softmax_lse)
32-{32+{
33- int64_t tSize = 0;33+ int64_t tSize = 0;
34- int64_t nSize = 0;34+ int64_t nSize = 0;
35- int64_t dSize = 0;35+ int64_t dSize = 0;
36- int64_t sSize = 0;36+ int64_t sSize = 0;
37- int64_t bSize = 0;37+ int64_t bSize = 0;
38- at::SmallVector<int64_t, MAX_DIM_SIZE> attentionOutSize;38+ at::SmallVector<int64_t, MAX_DIM_SIZE> attentionOutSize;
39- at::SmallVector<int64_t, MAX_DIM_SIZE> softmaxOutSize;39+ at::SmallVector<int64_t, MAX_DIM_SIZE> softmaxOutSize;
40- if (layout_q == "TND") {40+ if (layout_q == "TND") {
41- tSize = q.size(0);41+ tSize = q.size(0);
42- nSize = q.size(1);42+ nSize = q.size(1);
43- dSize = q.size(2);43+ dSize = q.size(2);
44- } else if (layout_q == "BSND") {44+ } else if (layout_q == "BSND") {
45- bSize = q.size(0);45+ bSize = q.size(0);
46- sSize = q.size(1);46+ sSize = q.size(1);
47- nSize = q.size(2);47+ nSize = q.size(2);
48- dSize = q.size(3);48+ dSize = q.size(3);
49- } else {49+ } else {
50- bSize = q.size(0);50+ bSize = q.size(0);
51- nSize = q.size(1);51+ nSize = q.size(1);
52- sSize = q.size(2);52+ sSize = q.size(2);
53- dSize = q.size(3);53+ dSize = q.size(3);
54- }54+ }
55- if (return_softmax_lse) {55+ if (return_softmax_lse) {
56- if (q.dim() == DIM_THREE) {56+ if (q.dim() == DIM_THREE) {
57- softmaxOutSize = {nSize, tSize};57+ softmaxOutSize = {nSize, tSize};
58- } else {58+ } else {
59- softmaxOutSize = {bSize, nSize, sSize};59+ softmaxOutSize = {bSize, nSize, sSize};
60- }60+ }
61- } else {61+ } else {
62- softmaxOutSize = {0};62+ softmaxOutSize = {0};
63- }63+ }
64- at::Tensor softmaxLse = at::empty(softmaxOutSize, q.options().dtype(at::kFloat));64+ at::Tensor softmaxLse = at::empty(softmaxOutSize, q.options().dtype(at::kFloat));
65- 65+ 
66- if (layout_out == "TND") {66+ if (layout_out == "TND") {
67- attentionOutSize = {tSize, nSize, dSize};67+ attentionOutSize = {tSize, nSize, dSize};
68- } else if (layout_out == "BNSD") {68+ } else if (layout_out == "BNSD") {
69- attentionOutSize = {bSize, nSize, sSize, dSize};69+ attentionOutSize = {bSize, nSize, sSize, dSize};
70- } else {70+ } else {
71- attentionOutSize = {bSize, sSize, nSize, dSize};71+ attentionOutSize = {bSize, sSize, nSize, dSize};
72- }72+ }
73- at::Tensor attentionOutput = at::empty(attentionOutSize, q.options().dtype(q.dtype()));73+ at::Tensor attentionOutput = at::empty(attentionOutSize, q.options().dtype(q.dtype()));
74- 74+ 
75- char *layout_q_ptr = const_cast<char *>(layout_q.c_str());75+ char *layout_q_ptr = const_cast<char *>(layout_q.c_str());
76- char *layout_kv_ptr = const_cast<char *>(layout_kv.c_str());76+ char *layout_kv_ptr = const_cast<char *>(layout_kv.c_str());
77- char *layout_out_ptr = const_cast<char *>(layout_out.c_str());77+ char *layout_out_ptr = const_cast<char *>(layout_out.c_str());
78- 78+ 
79- ACLNN_CMD(aclnnFlashAttn, q, k, v, block_table, cu_seqlens_q, cu_seqlens_kv, seqused_q, seqused_kv, sinks, attn_mask,79+ ACLNN_CMD(aclnnFlashAttn, q, k, v, block_table, cu_seqlens_q, cu_seqlens_kv, seqused_q, seqused_kv, sinks, attn_mask,
80- metadata, softmax_scale, mask_mode, win_left, win_right, max_seqlen_q, max_seqlen_kv, layout_q_ptr,80+ metadata, softmax_scale, mask_mode, win_left, win_right, max_seqlen_q, max_seqlen_kv, layout_q_ptr,
lilongwei_hw
lilongwei_hwlilongwei_hw6月16日
已过期

torch_extension 目录整改涉及公共头/宏迁移。请确认所有 OpBuilder 仍指向新路径且 ACLNN_CMD/ConvertType 行为一致,并跑 FIA/moe/mhc 冒烟 ST。

likedislike
81- layout_kv_ptr, layout_out_ptr, return_softmax_lse, attentionOutput, softmaxLse);81+ layout_kv_ptr, layout_out_ptr, return_softmax_lse, attentionOutput, softmaxLse);
82- 82+ 
83- return std::tuple<at::Tensor, at::Tensor>(attentionOutput, softmaxLse);83+ return std::tuple<at::Tensor, at::Tensor>(attentionOutput, softmaxLse);
84-}84+}
85-// Bind the C++ function to Python module85+// Bind the C++ function to Python module
86-PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)86+PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
87-{87+{
88- m.def("npu_flash_attn", &npu_flash_attn, "flash_attn");88+ m.def("npu_flash_attn", &npu_flash_attn, "flash_attn");
89-}89+}
90} // namespace op_api90} // namespace op_api
Rtorch_extension/npu_ops_transformer/ops/csrc/flash_attn_metadata.cpptorch_extension/cann_ops_transformer/ops/csrc/flash_attn_metadata.cpp+40-40
@@ -1,40 +1,40 @@
1-/**1+/**
2- * Copyright (c) 2026 Huawei Technologies Co., Ltd.2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3- * This program is free software, you can redistribute it and/or modify it under the terms and conditions of3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4- * CANN Open Software License Agreement Version 2.0 (the "License").4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5- * Please refer to the License for details. You may not use this file except in compliance with the License.5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6- * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7- * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8- * See LICENSE in the root of the software repository for the full text of the License.8+ * See LICENSE in the root of the software repository for the full text of the License.
9- */9+ */
10- 10+ 
11-/*!11+/*!
12- * \file flash_attn_metadata.cpp12+ * \file flash_attn_metadata.cpp
13- * \brief13+ * \brief
14- */14+ */
15- 15+ 
16-#include <torch/extension.h>16+#include <torch/extension.h>
17-#include "aclnn_common.h"17+#include "aclnn_common.h"
18- 18+ 
19-namespace op_api {19+namespace op_api {
20-using npu_utils = at_npu::native::NpuUtils;20+using npu_utils = at_npu::native::NpuUtils;
21- 21+ 
22-at::Tensor npu_flash_attn_metadata(const c10::optional<at::Tensor> &cu_seqlens_q,22+at::Tensor npu_flash_attn_metadata(const c10::optional<at::Tensor> &cu_seqlens_q,
23- const c10::optional<at::Tensor> &cu_seqlens_kv,23+ const c10::optional<at::Tensor> &cu_seqlens_kv,
24- const c10::optional<at::Tensor> &seqused_q,24+ const c10::optional<at::Tensor> &seqused_q,
25- const c10::optional<at::Tensor> &seqused_kv, int64_t num_heads_q,25+ const c10::optional<at::Tensor> &seqused_kv, int64_t num_heads_q,
26- int64_t num_heads_kv, int64_t head_dim, int64_t batch_size, int64_t max_seqlen_q,26+ int64_t num_heads_kv, int64_t head_dim, int64_t batch_size, int64_t max_seqlen_q,
27- int64_t max_seqlen_kv, int64_t mask_mode, int64_t win_left, int64_t win_right,27+ int64_t max_seqlen_kv, int64_t mask_mode, int64_t win_left, int64_t win_right,
28- std::string layout_q, std::string layout_kv, std::string layout_out, const at::Tensor &output)28+ std::string layout_q, std::string layout_kv, std::string layout_out, const at::Tensor &output)
29-{29+{
30- ACLNN_CMD(aclnnFlashAttnMetadata, cu_seqlens_q, cu_seqlens_kv, seqused_q, seqused_kv, batch_size, max_seqlen_q,30+ ACLNN_CMD(aclnnFlashAttnMetadata, cu_seqlens_q, cu_seqlens_kv, seqused_q, seqused_kv, batch_size, max_seqlen_q,
31- max_seqlen_kv, num_heads_q, num_heads_kv, head_dim, mask_mode, win_left, win_right, layout_q, layout_kv,31+ max_seqlen_kv, num_heads_q, num_heads_kv, head_dim, mask_mode, win_left, win_right, layout_q, layout_kv,
32- layout_out, output);32+ layout_out, output);
33- return output;33+ return output;
34-}34+}
35- 35+ 
36-PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)36+PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
37-{37+{
38- m.def("npu_flash_attn_metadata", &npu_flash_attn_metadata, "npu_flash_attn_metadata");38+ m.def("npu_flash_attn_metadata", &npu_flash_attn_metadata, "npu_flash_attn_metadata");
39-}39+}
40-} // namespace op_api40+} // namespace op_api
Rtorch_extension/npu_ops_transformer/ops/csrc/lightning_indexer_v2.cpptorch_extension/cann_ops_transformer/ops/csrc/lightning_indexer_v2.cpp+147-106
@@ -1,107 +1,148 @@
1-/**1+/**
2-* Copyright (c) 2026 Huawei Technologies Co., Ltd.2+* Copyright (c) 2026 Huawei Technologies Co., Ltd.
3-* This program is free software, you can redistribute it and/or modify it under the terms and conditions of3+* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4-* CANN Open Software License Agreement Version 2.0 (the "License").4+* CANN Open Software License Agreement Version 2.0 (the "License").
5-* Please refer to the License for details. You may not use this file except in compliance with the License.5+* Please refer to the License for details. You may not use this file except in compliance with the License.
6-* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,6+* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7-* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.7+* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8-* See LICENSE in the root of the software repository for the full text of the License.8+* See LICENSE in the root of the software repository for the full text of the License.
9-*/9+*/
10- 10+ 
11-/*!11+/*!
12-* \file lightning_indexer_v2.cpp12+* \file lightning_indexer_v2.cpp
13-* \brief13+* \brief
14-*/14+*/
15- 15+ 
16-#include <torch/extension.h>16+#include <torch/extension.h>
17-#include "aclnn_common.h"17+#include "aclnn_common.h"
18- 18+ 
19-namespace op_api {19+namespace op_api {
20-using namespace at_npu::native;20+using namespace at_npu::native;
21- 21+ 
22-// npu tensor max size22+// npu tensor max size
23-const int SIZE = 8;23+const int SIZE = 8;
24-const int DIM_0 = 0;24+const int DIM_0 = 0;
25-const int DIM_1 = 1;25+const int DIM_1 = 1;
26-const int DIM_2 = 2;26+const int DIM_2 = 2;
27-const int DIM_3 = 3;27+const int DIM_3 = 3;
28- 28+ 
29-// 工具函数,推导输出shape29+constexpr int64_t LI_V2_METADATA_SIZE = 1024;
30-std::tuple<at::Tensor, at::Tensor> construct_lightning_indexer_v2_output_tensor(const at::Tensor& query,30+ 
31- const at::Tensor& key,31+at::Tensor lightning_indexer_metadata(
32- int64_t topk,32+ int64_t num_heads_q, int64_t num_heads_k, int64_t head_dim, int64_t topk,
33- std::string query_layout_str,33+ const c10::optional<at::Tensor> &cu_seqlens_q, const c10::optional<at::Tensor> &cu_seqlens_k,
34- std::string key_layout_str,34+ const c10::optional<at::Tensor> &seqused_q, const c10::optional<at::Tensor> &seqused_k,
35- bool return_value)35+ const c10::optional<at::Tensor> &cmp_residual_k, int64_t batch_size, int64_t max_seqlen_q, int64_t max_seqlen_k,
36-{36+ c10::string_view layout_q, c10::string_view layout_k, int64_t mask_mode, int64_t cmp_ratio)
37- at::SmallVector<int64_t, SIZE> output_size;37+{
38- for (size_t i = 0; i < query.sizes().size(); i++) {38+ at::Device output_device = at::Device(std::string("npu"));
39- TORCH_CHECK(query.size(i) > 0, "All values within query's shape should be greater "39+ if (cu_seqlens_q.has_value()) {
40- "than 0, but shape[", i, "] is ", query.size(i));40+ output_device = cu_seqlens_q.value().device();
41- }41+ } else if (cu_seqlens_k.has_value()) {
42- for (size_t i = 0; i < key.sizes().size(); i++) {42+ output_device = cu_seqlens_k.value().device();
43- TORCH_CHECK(key.size(i) > 0, "All values within key's shape should be greater "43+ } else if (seqused_q.has_value()) {
44- "than 0, but shape[", i, "] is ", key.size(i));44+ output_device = seqused_q.value().device();
45- }45+ } else if (seqused_k.has_value()) {
46- TORCH_CHECK(topk > 0, "topk should be greater than 0, but now is ", topk);46+ output_device = seqused_k.value().device();
47- int64_t keyHeadNum = (key_layout_str == "TND")? key.size(DIM_1) : key.size(DIM_2);47+ } else if (cmp_residual_k.has_value()) {
48- if (query_layout_str == "BSND") {48+ output_device = cmp_residual_k.value().device();
49- output_size = {query.size(DIM_0), query.size(DIM_1), keyHeadNum, topk};49+ }
50- } else {50+ 
51- int n_dim_index = 0;51+ at::Tensor output = torch::empty({LI_V2_METADATA_SIZE}, torch::dtype(torch::kInt32).device(output_device));
52- n_dim_index = (key_layout_str == "TND") ? DIM_1 : DIM_2;52+ auto cu_seqlens_q_val = get_valid_tensor(cu_seqlens_q, output_device);
53- output_size = {query.size(DIM_0), key.size(n_dim_index), topk};53+ auto cu_seqlens_k_val = get_valid_tensor(cu_seqlens_k, output_device);
54- }54+ auto seqused_q_val = get_valid_tensor(seqused_q, output_device);
55- at::Tensor sparse_indices_out = at::empty(output_size, query.options().dtype(at::kInt));55+ auto seqused_k_val = get_valid_tensor(seqused_k, output_device);
56- at::Tensor sparse_values_out;56+ auto cmp_residual_k_val = get_valid_tensor(cmp_residual_k, output_device);
57- if (return_value) {57+ 
58- sparse_values_out = at::empty(output_size, query.options().dtype(at::kFloat));58+ std::string layout_q_str = std::string(layout_q);
59- } else {59+ std::string layout_k_str = std::string(layout_k);
60- sparse_values_out = at::empty({0}, query.options().dtype(at::kFloat));60+ char *layout_q_ptr = const_cast<char *>(layout_q_str.c_str());
61- }61+ char *layout_k_ptr = const_cast<char *>(layout_k_str.c_str());
62- 62+ 
63- return std::tuple<at::Tensor, at::Tensor>(sparse_indices_out, sparse_values_out);63+ ACLNN_CMD(aclnnLightningIndexerV2Metadata, cu_seqlens_q_val, cu_seqlens_k_val, seqused_q_val, seqused_k_val,
64-}64+ cmp_residual_k_val, num_heads_q, num_heads_k, head_dim, topk, batch_size, max_seqlen_q, max_seqlen_k,
65- 65+ layout_q_ptr, layout_k_ptr, mask_mode, cmp_ratio, output);
66-std::tuple<at::Tensor, at::Tensor> npu_lightning_indexer_v2(66+ return output;
67- const at::Tensor &q, const at::Tensor &k, const at::Tensor &w,67+}
68- int64_t topk,68+ 
69- const c10::optional<at::Tensor> &cu_seqlens_q,69+// 工具函数,推导输出shape
70- const c10::optional<at::Tensor> &cu_seqlens_k,70+std::tuple<at::Tensor, at::Tensor> construct_lightning_indexer_v2_output_tensor(const at::Tensor& query,
71- const c10::optional<at::Tensor> &seqused_q,71+ const at::Tensor& key,
72- const c10::optional<at::Tensor> &seqused_k,72+ int64_t topk,
73- const c10::optional<at::Tensor> &cmpResidualK,73+ std::string query_layout_str,
74- const c10::optional<at::Tensor> &block_table,74+ std::string key_layout_str,
75- const c10::optional<at::Tensor> &output_idx_offset,75+ bool return_value)
76- const c10::optional<at::Tensor> &metadata,76+{
77- int64_t max_seqlen_q,77+ at::SmallVector<int64_t, SIZE> output_size;
78- c10::string_view layout_q, c10::string_view layout_k,78+ for (size_t i = 0; i < query.sizes().size(); i++) {
79- int64_t mask_mode, int64_t cmp_ratio, int64_t return_value)79+ TORCH_CHECK(query.size(i) > 0, "All values within query's shape should be greater "
80-{80+ "than 0, but shape[", i, "] is ", query.size(i));
81- TORCH_CHECK(q.numel() > 0, "Tensor q is empty.")81+ }
82- TORCH_CHECK(k.numel() > 0, "Tensor k is empty.")82+ for (size_t i = 0; i < key.sizes().size(); i++) {
83- 83+ TORCH_CHECK(key.size(i) > 0, "All values within key's shape should be greater "
84- std::string query_layout_str = std::string(layout_q);84+ "than 0, but shape[", i, "] is ", key.size(i));
85- std::string key_layout_str = std::string(layout_k);85+ }
86- 86+ TORCH_CHECK(topk > 0, "topk should be greater than 0, but now is ", topk);
87- // construct the output tensor87+ int64_t keyHeadNum = (key_layout_str == "TND")? key.size(DIM_1) : key.size(DIM_2);
88- std::tuple<at::Tensor, at::Tensor> lightning_indexer_v2_output88+ if (query_layout_str == "BSND") {
89- = construct_lightning_indexer_v2_output_tensor(q, k, topk, query_layout_str,89+ output_size = {query.size(DIM_0), query.size(DIM_1), keyHeadNum, topk};
90- key_layout_str, return_value);90+ } else {
91- at::Tensor sparse_indices_out = std::get<0>(lightning_indexer_v2_output);91+ int n_dim_index = 0;
92- at::Tensor sparse_values_out = std::get<1>(lightning_indexer_v2_output);92+ n_dim_index = (key_layout_str == "TND") ? DIM_1 : DIM_2;
93- // convert str93+ output_size = {query.size(DIM_0), key.size(n_dim_index), topk};
94- char *query_layout_ptr = const_cast<char *>(query_layout_str.c_str());94+ }
95- char *key_layout_ptr = const_cast<char *>(key_layout_str.c_str());95+ at::Tensor sparse_indices_out = at::empty(output_size, query.options().dtype(at::kInt));
96- 96+ at::Tensor sparse_values_out;
97- ACLNN_CMD(aclnnLightningIndexerV2, q, k, w, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k, cmpResidualK,97+ if (return_value) {
98- block_table, output_idx_offset, metadata, topk, max_seqlen_q, query_layout_ptr, key_layout_ptr,98+ sparse_values_out = at::empty(output_size, query.options().dtype(at::kFloat));
99- mask_mode, cmp_ratio, return_value, sparse_indices_out, sparse_values_out);99+ } else {
100- return std::tuple<at::Tensor, at::Tensor>(sparse_indices_out, sparse_values_out);100+ sparse_values_out = at::empty({0}, query.options().dtype(at::kFloat));
101-}101+ }
102-// Bind the C++ function to Python module102+ 
103-PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)103+ return std::tuple<at::Tensor, at::Tensor>(sparse_indices_out, sparse_values_out);
104-{104+}
105- m.def("npu_lightning_indexer_v2", &npu_lightning_indexer_v2, "lightning_indexer_v2");105+ 
106-}106+std::tuple<at::Tensor, at::Tensor> lightning_indexer_v2(
107+ const at::Tensor &q, const at::Tensor &k, const at::Tensor &w,
108+ int64_t topk,
109+ const c10::optional<at::Tensor> &cu_seqlens_q,
110+ const c10::optional<at::Tensor> &cu_seqlens_k,
111+ const c10::optional<at::Tensor> &seqused_q,
112+ const c10::optional<at::Tensor> &seqused_k,
113+ const c10::optional<at::Tensor> &cmpResidualK,
114+ const c10::optional<at::Tensor> &block_table,
115+ const c10::optional<at::Tensor> &output_idx_offset,
116+ const c10::optional<at::Tensor> &metadata,
117+ int64_t max_seqlen_q,
118+ c10::string_view layout_q, c10::string_view layout_k,
119+ int64_t mask_mode, int64_t cmp_ratio, int64_t return_value)
120+{
121+ TORCH_CHECK(q.numel() > 0, "Tensor q is empty.")
122+ TORCH_CHECK(k.numel() > 0, "Tensor k is empty.")
123+ 
124+ std::string query_layout_str = std::string(layout_q);
125+ std::string key_layout_str = std::string(layout_k);
126+ 
127+ // construct the output tensor
128+ std::tuple<at::Tensor, at::Tensor> lightning_indexer_v2_output
129+ = construct_lightning_indexer_v2_output_tensor(q, k, topk, query_layout_str,
130+ key_layout_str, return_value);
131+ at::Tensor sparse_indices_out = std::get<0>(lightning_indexer_v2_output);
132+ at::Tensor sparse_values_out = std::get<1>(lightning_indexer_v2_output);
133+ // convert str
134+ char *query_layout_ptr = const_cast<char *>(query_layout_str.c_str());
135+ char *key_layout_ptr = const_cast<char *>(key_layout_str.c_str());
136+ 
137+ ACLNN_CMD(aclnnLightningIndexerV2, q, k, w, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k, cmpResidualK,
138+ block_table, output_idx_offset, metadata, topk, max_seqlen_q, query_layout_ptr, key_layout_ptr,
139+ mask_mode, cmp_ratio, return_value, sparse_indices_out, sparse_values_out);
140+ return std::tuple<at::Tensor, at::Tensor>(sparse_indices_out, sparse_values_out);
141+}
142+// Bind the C++ function to Python module
143+PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
144+{
145+ m.def("lightning_indexer_metadata", &lightning_indexer_metadata, "lightning_indexer_metadata");
146+ m.def("lightning_indexer_v2", &lightning_indexer_v2, "lightning_indexer_v2");
147+}
107} // namespace op_api148} // namespace op_api
Rtorch_extension/npu_ops_transformer/ops/csrc/mega_moe.cpptorch_extension/cann_ops_transformer/ops/csrc/mega_moe.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/mhc_post.cpptorch_extension/cann_ops_transformer/ops/csrc/mhc_post.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/mhc_post_backward.cpptorch_extension/cann_ops_transformer/ops/csrc/mhc_post_backward.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/mhc_pre_sinkhorn.cpptorch_extension/cann_ops_transformer/ops/csrc/mhc_pre_sinkhorn.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/mhc_pre_sinkhorn_backward.cpptorch_extension/cann_ops_transformer/ops/csrc/mhc_pre_sinkhorn_backward.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/mixed_quant_sparse_flash_mla.cpptorch_extension/cann_ops_transformer/ops/csrc/mixed_quant_sparse_flash_mla.cpp+64-2
@@ -25,6 +25,66 @@ const int DIM_2 = 2;
25const int DIM_3 = 3;25const int DIM_3 = 3;
26const int DIM_4 = 4;26const int DIM_4 = 4;
27 27 
28+constexpr int64_t MQSMLA_METADATA_SIZE = 1024;
29+ 
30+at::Tensor mixed_quant_sparse_flash_mla_metadata(
31+ int64_t num_heads_q, int64_t num_heads_kv, int64_t head_dim, int64_t quant_mode,
32+ const c10::optional<at::Tensor> &cu_seqlens_q, const c10::optional<at::Tensor> &cu_seqlens_ori_kv,
33+ const c10::optional<at::Tensor> &cu_seqlens_cmp_kv, const c10::optional<at::Tensor> &seqused_q,
34+ const c10::optional<at::Tensor> &seqused_ori_kv, const c10::optional<at::Tensor> &seqused_cmp_kv,
35+ const c10::optional<at::Tensor> &cmp_residual_kv, const c10::optional<at::Tensor> &ori_topk_length,
36+ const c10::optional<at::Tensor> &cmp_topk_length, int64_t batch_size, int64_t max_seqlen_q,
37+ int64_t max_seqlen_ori_kv, int64_t max_seqlen_cmp_kv, int64_t ori_topk, int64_t cmp_topk, int64_t rope_head_dim,
38+ int64_t cmp_ratio, int64_t ori_mask_mode, int64_t cmp_mask_mode, int64_t ori_win_left, int64_t ori_win_right,
39+ c10::string_view layout_q, c10::string_view layout_kv, bool has_ori_kv, bool has_cmp_kv)
40+{
41+ at::Device output_device = at::Device(std::string("npu"));
42+ if (cu_seqlens_q.has_value()) {
43+ output_device = cu_seqlens_q.value().device();
44+ } else if (cu_seqlens_ori_kv.has_value()) {
45+ output_device = cu_seqlens_ori_kv.value().device();
46+ } else if (cu_seqlens_cmp_kv.has_value()) {
47+ output_device = cu_seqlens_cmp_kv.value().device();
48+ } else if (seqused_q.has_value()) {
49+ output_device = seqused_q.value().device();
50+ } else if (seqused_ori_kv.has_value()) {
51+ output_device = seqused_ori_kv.value().device();
52+ } else if (seqused_cmp_kv.has_value()) {
53+ output_device = seqused_cmp_kv.value().device();
54+ } else if (cmp_residual_kv.has_value()) {
55+ output_device = cmp_residual_kv.value().device();
56+ } else if (ori_topk_length.has_value()) {
57+ output_device = ori_topk_length.value().device();
58+ } else if (cmp_topk_length.has_value()) {
59+ output_device = cmp_topk_length.value().device();
60+ }
61+ 
62+ at::Tensor output = torch::empty({MQSMLA_METADATA_SIZE}, torch::dtype(torch::kInt32).device(output_device));
63+ auto cu_seqlens_q_val = get_valid_tensor(cu_seqlens_q, output_device);
64+ auto cu_seqlens_ori_kv_val = get_valid_tensor(cu_seqlens_ori_kv, output_device);
65+ auto cu_seqlens_cmp_kv_val = get_valid_tensor(cu_seqlens_cmp_kv, output_device);
66+ auto seqused_q_val = get_valid_tensor(seqused_q, output_device);
67+ auto seqused_ori_kv_val = get_valid_tensor(seqused_ori_kv, output_device);
68+ auto seqused_cmp_kv_val = get_valid_tensor(seqused_cmp_kv, output_device);
69+ auto cmp_residual_kv_val = get_valid_tensor(cmp_residual_kv, output_device);
70+ auto ori_topk_length_val = get_valid_tensor(ori_topk_length, output_device);
71+ auto cmp_topk_length_val = get_valid_tensor(cmp_topk_length, output_device);
72+ 
73+ // convert str
74+ std::string layout_q_str = std::string(layout_q);
75+ std::string layout_kv_str = std::string(layout_kv);
76+ char *layout_q_ptr = const_cast<char *>(layout_q_str.c_str());
77+ char *layout_kv_ptr = const_cast<char *>(layout_kv_str.c_str());
78+ 
79+ ACLNN_CMD(aclnnMixedQuantSparseFlashMlaMetadata, cu_seqlens_q_val, cu_seqlens_ori_kv_val, cu_seqlens_cmp_kv_val,
80+ seqused_q_val, seqused_ori_kv_val, seqused_cmp_kv_val, cmp_residual_kv_val, ori_topk_length_val,
81+ cmp_topk_length_val, num_heads_q, num_heads_kv, head_dim, quant_mode, batch_size, max_seqlen_q,
82+ max_seqlen_ori_kv, max_seqlen_cmp_kv, ori_topk, cmp_topk, rope_head_dim, cmp_ratio, ori_mask_mode,
83+ cmp_mask_mode, ori_win_left, ori_win_right, layout_q_ptr, layout_kv_ptr, has_ori_kv, has_cmp_kv,
84+ output);
85+ return output;
86+}
87+ 
28std::tuple<at::Tensor, at::Tensor> construct_mixed_quant_sparse_flash_mla_atten_out_tensor(88std::tuple<at::Tensor, at::Tensor> construct_mixed_quant_sparse_flash_mla_atten_out_tensor(
29 const at::Tensor& q, const at::Tensor& ori_kv, std::string layout_q_str,89 const at::Tensor& q, const at::Tensor& ori_kv, std::string layout_q_str,
30 std::string layout_kv_str, const uint64_t &rope_head_dim, bool return_softmax_lse)90 std::string layout_kv_str, const uint64_t &rope_head_dim, bool return_softmax_lse)
@@ -89,7 +149,7 @@ std::tuple<at::Tensor, at::Tensor> construct_mixed_quant_sparse_flash_mla_atten_
89 return std::tuple<at::Tensor, at::Tensor>(atten_out, softmax_lse);149 return std::tuple<at::Tensor, at::Tensor>(atten_out, softmax_lse);
90}150}
91 151 
92-std::tuple<at::Tensor, at::Tensor> npu_mixed_quant_sparse_flash_mla(152+std::tuple<at::Tensor, at::Tensor> mixed_quant_sparse_flash_mla(
93 const at::Tensor &q,153 const at::Tensor &q,
94 const c10::optional<at::Tensor> &ori_kv, const c10::optional<at::Tensor> &cmp_kv,154 const c10::optional<at::Tensor> &ori_kv, const c10::optional<at::Tensor> &cmp_kv,
95 const c10::optional<at::Tensor> &ori_sparse_indices, const c10::optional<at::Tensor> &cmp_sparse_indices,155 const c10::optional<at::Tensor> &ori_sparse_indices, const c10::optional<at::Tensor> &cmp_sparse_indices,
@@ -158,6 +218,8 @@ std::tuple<at::Tensor, at::Tensor> npu_mixed_quant_sparse_flash_mla(
158 218 
159PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)219PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
160{220{
161- m.def("npu_mixed_quant_sparse_flash_mla", &npu_mixed_quant_sparse_flash_mla, "npu_mixed_quant_sparse_flash_mla");221+ m.def("mixed_quant_sparse_flash_mla_metadata", &mixed_quant_sparse_flash_mla_metadata,
222+ "mixed_quant_sparse_flash_mla_metadata");
223+ m.def("mixed_quant_sparse_flash_mla", &mixed_quant_sparse_flash_mla, "mixed_quant_sparse_flash_mla");
162}224}
163} // namespace op_api225} // namespace op_api
Rtorch_extension/npu_ops_transformer/ops/csrc/moe_distribute_combine.cpptorch_extension/cann_ops_transformer/ops/csrc/moe_distribute_combine.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/moe_distribute_dispatch.cpptorch_extension/cann_ops_transformer/ops/csrc/moe_distribute_dispatch.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/npu_sparse_lightning_indexer_kl_loss_grad.cpptorch_extension/cann_ops_transformer/ops/csrc/npu_sparse_lightning_indexer_kl_loss_grad.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/npu_sparse_lightning_indexer_kl_loss_grad_metadata.cpptorch_extension/cann_ops_transformer/ops/csrc/npu_sparse_lightning_indexer_kl_loss_grad_metadata.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/quant_lightning_indexer_v2.cpptorch_extension/cann_ops_transformer/ops/csrc/quant_lightning_indexer_v2.cpp+44-2
@@ -26,6 +26,46 @@ const int DIM_1 = 1;
26const int DIM_2 = 2;26const int DIM_2 = 2;
27const int DIM_3 = 3;27const int DIM_3 = 3;
28 28 
29+constexpr int64_t QLI_V2_METADATA_SIZE = 1024;
30+ 
31+at::Tensor quant_lightning_indexer_metadata(
32+ int64_t num_heads_q, int64_t num_heads_k, int64_t head_dim, int64_t topk, int64_t quant_mode,
33+ const c10::optional<at::Tensor> &cu_seqlens_q, const c10::optional<at::Tensor> &cu_seqlens_k,
34+ const c10::optional<at::Tensor> &seqused_q, const c10::optional<at::Tensor> &seqused_k,
35+ const c10::optional<at::Tensor> &cmp_residual_k, int64_t batch_size, int64_t max_seqlen_q, int64_t max_seqlen_k,
36+ c10::string_view layout_q, c10::string_view layout_k, int64_t mask_mode, int64_t cmp_ratio)
37+{
38+ at::Device output_device = at::Device(std::string("npu"));
39+ if (cu_seqlens_q.has_value()) {
40+ output_device = cu_seqlens_q.value().device();
41+ } else if (cu_seqlens_k.has_value()) {
42+ output_device = cu_seqlens_k.value().device();
43+ } else if (seqused_q.has_value()) {
44+ output_device = seqused_q.value().device();
45+ } else if (seqused_k.has_value()) {
46+ output_device = seqused_k.value().device();
47+ } else if (cmp_residual_k.has_value()) {
48+ output_device = cmp_residual_k.value().device();
49+ }
50+ 
51+ at::Tensor output = torch::empty({QLI_V2_METADATA_SIZE}, torch::dtype(torch::kInt32).device(output_device));
52+ auto cu_seqlens_q_val = get_valid_tensor(cu_seqlens_q, output_device);
53+ auto cu_seqlens_k_val = get_valid_tensor(cu_seqlens_k, output_device);
54+ auto seqused_q_val = get_valid_tensor(seqused_q, output_device);
55+ auto seqused_k_val = get_valid_tensor(seqused_k, output_device);
56+ auto cmp_residual_k_val = get_valid_tensor(cmp_residual_k, output_device);
57+ 
58+ std::string layout_q_str = std::string(layout_q);
59+ std::string layout_k_str = std::string(layout_k);
60+ char *layout_q_ptr = const_cast<char *>(layout_q_str.c_str());
61+ char *layout_k_ptr = const_cast<char *>(layout_k_str.c_str());
62+ 
63+ ACLNN_CMD(aclnnQuantLightningIndexerV2Metadata, cu_seqlens_q_val, cu_seqlens_k_val, seqused_q_val, seqused_k_val,
64+ cmp_residual_k_val, num_heads_q, num_heads_k, head_dim, topk, quant_mode, batch_size, max_seqlen_q,
65+ max_seqlen_k, layout_q_ptr, layout_k_ptr, mask_mode, cmp_ratio, output);
66+ return output;
67+}
68+ 
29// 工具函数,推导输出shape69// 工具函数,推导输出shape
30std::tuple<at::Tensor, at::Tensor> construct_quant_lightning_indexer_output_tensor(70std::tuple<at::Tensor, at::Tensor> construct_quant_lightning_indexer_output_tensor(
31 const at::Tensor& query, const at::Tensor& key,71 const at::Tensor& query, const at::Tensor& key,
@@ -61,7 +101,7 @@ std::tuple<at::Tensor, at::Tensor> construct_quant_lightning_indexer_output_tens
61 return std::tuple<at::Tensor, at::Tensor>(sparse_indices_out, sparse_values_out);101 return std::tuple<at::Tensor, at::Tensor>(sparse_indices_out, sparse_values_out);
62}102}
63 103 
64-std::tuple<at::Tensor, at::Tensor> npu_quant_lightning_indexer_v2(104+std::tuple<at::Tensor, at::Tensor> quant_lightning_indexer_v2(
65 const at::Tensor &query, const at::Tensor &key, const at::Tensor &weights,105 const at::Tensor &query, const at::Tensor &key, const at::Tensor &weights,
66 const at::Tensor &query_dequant_scale, const at::Tensor &key_dequant_scale,106 const at::Tensor &query_dequant_scale, const at::Tensor &key_dequant_scale,
67 int64_t topk, int64_t quant_mode,107 int64_t topk, int64_t quant_mode,
@@ -102,6 +142,8 @@ std::tuple<at::Tensor, at::Tensor> npu_quant_lightning_indexer_v2(
102// Bind the C++ function to Python module142// Bind the C++ function to Python module
103PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)143PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
104{144{
105- m.def("npu_quant_lightning_indexer_v2", &npu_quant_lightning_indexer_v2, "quant_lightning_indexer_v2");145+ m.def("quant_lightning_indexer_metadata", &quant_lightning_indexer_metadata,
146+ "quant_lightning_indexer_metadata");
147+ m.def("quant_lightning_indexer_v2", &quant_lightning_indexer_v2, "quant_lightning_indexer_v2");
106}148}
107} // namespace op_api149} // namespace op_api
@@ -0,0 +1,207 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+/*!
11+ * \file sparse_flash_mla.cpp
12+ * \brief
13+ */
14+ 
15+#include <torch/extension.h>
16+#include "aclnn_common.h"
17+ 
18+namespace op_api {
19+using namespace at_npu::native;
20+using npu_preparation = at_npu::native::OpPreparation;
21+ 
22+// npu tensor max size
23+const int SIZE = 8;
24+const int DIM_0 = 0;
25+const int DIM_1 = 1;
26+const int DIM_2 = 2;
27+const int DIM_3 = 3;
28+const int DIM_4 = 4;
29+ 
30+constexpr int64_t SMLA_METADATA_SIZE = 1024;
31+ 
32+const c10::optional<at::Tensor> smla_get_valid_tensor(const c10::optional<at::Tensor> &tensor_opt, at::Device device)
33+{
34+ return tensor_opt.has_value() ? tensor_opt : torch::empty({0}, torch::dtype(torch::kInt32).device(device));
35+};
36+ 
37+at::Tensor sparse_flash_mla_metadata(
38+ int64_t num_heads_q, int64_t num_heads_kv, int64_t head_dim, const c10::optional<at::Tensor> &cu_seqlens_q,
39+ const c10::optional<at::Tensor> &cu_seqlens_ori_kv, const c10::optional<at::Tensor> &cu_seqlens_cmp_kv,
40+ const c10::optional<at::Tensor> &seqused_q, const c10::optional<at::Tensor> &seqused_ori_kv,
41+ const c10::optional<at::Tensor> &seqused_cmp_kv, const c10::optional<at::Tensor> &cmp_residual_kv,
42+ const c10::optional<at::Tensor> &ori_topk_length, const c10::optional<at::Tensor> &cmp_topk_length,
43+ int64_t batch_size, int64_t max_seqlen_q, int64_t max_seqlen_ori_kv, int64_t max_seqlen_cmp_kv, int64_t ori_topk,
44+ int64_t cmp_topk, int64_t cmp_ratio, int64_t ori_mask_mode, int64_t cmp_mask_mode, int64_t ori_win_left,
45+ int64_t ori_win_right, c10::string_view layout_q, c10::string_view layout_kv, bool has_ori_kv, bool has_cmp_kv)
46+{
47+ at::Device output_device = at::Device(std::string("npu"));
48+ if (cu_seqlens_q.has_value()) {
49+ output_device = cu_seqlens_q.value().device();
50+ } else if (cu_seqlens_ori_kv.has_value()) {
51+ output_device = cu_seqlens_ori_kv.value().device();
52+ } else if (cu_seqlens_cmp_kv.has_value()) {
53+ output_device = cu_seqlens_cmp_kv.value().device();
54+ } else if (seqused_q.has_value()) {
55+ output_device = seqused_q.value().device();
56+ } else if (seqused_ori_kv.has_value()) {
57+ output_device = seqused_ori_kv.value().device();
58+ } else if (seqused_cmp_kv.has_value()) {
59+ output_device = seqused_cmp_kv.value().device();
60+ } else if (cmp_residual_kv.has_value()) {
61+ output_device = cmp_residual_kv.value().device();
62+ } else if (ori_topk_length.has_value()) {
63+ output_device = ori_topk_length.value().device();
64+ } else if (cmp_topk_length.has_value()) {
65+ output_device = cmp_topk_length.value().device();
66+ }
67+ 
68+ at::Tensor output = torch::empty({SMLA_METADATA_SIZE}, torch::dtype(torch::kInt32).device(output_device));
69+ auto cu_seqlens_q_val = smla_get_valid_tensor(cu_seqlens_q, output_device);
70+ auto cu_seqlens_ori_kv_val = smla_get_valid_tensor(cu_seqlens_ori_kv, output_device);
71+ auto cu_seqlens_cmp_kv_val = smla_get_valid_tensor(cu_seqlens_cmp_kv, output_device);
72+ auto seqused_q_val = smla_get_valid_tensor(seqused_q, output_device);
73+ auto seqused_ori_kv_val = smla_get_valid_tensor(seqused_ori_kv, output_device);
74+ auto seqused_cmp_kv_val = smla_get_valid_tensor(seqused_cmp_kv, output_device);
75+ auto cmp_residual_kv_val = smla_get_valid_tensor(cmp_residual_kv, output_device);
76+ auto ori_topk_length_val = smla_get_valid_tensor(ori_topk_length, output_device);
77+ auto cmp_topk_length_val = smla_get_valid_tensor(cmp_topk_length, output_device);
78+ 
79+ // convert str
80+ std::string layout_q_str = std::string(layout_q);
81+ std::string layout_kv_str = std::string(layout_kv);
82+ char *layout_q_ptr = const_cast<char *>(layout_q_str.c_str());
83+ char *layout_kv_ptr = const_cast<char *>(layout_kv_str.c_str());
84+ 
85+ ACLNN_CMD(aclnnSparseFlashMlaMetadata, cu_seqlens_q_val, cu_seqlens_ori_kv_val, cu_seqlens_cmp_kv_val,
86+ seqused_q_val, seqused_ori_kv_val, seqused_cmp_kv_val, cmp_residual_kv_val, ori_topk_length_val,
87+ cmp_topk_length_val, num_heads_q, num_heads_kv, head_dim, batch_size, max_seqlen_q, max_seqlen_ori_kv,
88+ max_seqlen_cmp_kv, ori_topk, cmp_topk, cmp_ratio, ori_mask_mode, cmp_mask_mode, ori_win_left,
89+ ori_win_right, layout_q_ptr, layout_kv_ptr, has_ori_kv, has_cmp_kv, output);
90+ return output;
91+}
92+ 
93+std::tuple<at::Tensor, at::Tensor> construct_sparse_flash_mla_atten_out_tensor(
94+ const at::Tensor& q, const at::Tensor& ori_kv, std::string layout_q_str,
95+ std::string layout_kv_str, const uint64_t &rope_head_dim, bool return_softmax_lse)
96+{
97+ TORCH_CHECK(layout_q_str == "BSND" || layout_q_str == "TND",
98+ "The layout of query only support BSND and TND, but got ", layout_q_str);
99+ for (auto i = 0; i < q.sizes().size(); i++) {
100+ TORCH_CHECK(q.size(i) > 0, "All values within query's shape should be greater "
101+ "than 0, but shape[", i, "] is ", q.size(i));
102+ }
103+ at::SmallVector<int64_t, SIZE> atten_out_size;
104+ at::SmallVector<int64_t, SIZE> softmax_lse_size;
105+ if (layout_q_str == "BSND") {
106+ TORCH_CHECK(q.dim() == DIM_4,
107+ "When the layout of query is BSND, the query dimension must be 4, but got ", q.dim());
108+ atten_out_size = {q.size(DIM_0), q.size(DIM_1), q.size(DIM_2), q.size(DIM_3)};
109+ } else {
110+ TORCH_CHECK(q.dim() == DIM_3,
111+ "When the layout of query is TND, the query dimension must be 3, but got ", q.dim());
112+ atten_out_size = {q.size(DIM_0), q.size(DIM_1), q.size(DIM_2)};
113+ }
114+ at::Tensor atten_out = at::empty(atten_out_size, q.options().dtype(q.dtype()));
115+ at::Tensor softmax_lse;
116+ 
117+ if (return_softmax_lse) {
118+ if (layout_q_str == "BSND") {
119+ // 对齐 Python: [q.shape[0], ori_kv.shape[2], q.shape[1], q.shape[2] // ori_kv.shape[2]]
120+ softmax_lse_size = {
121+ q.size(DIM_0),
122+ ori_kv.size(DIM_2),
123+ q.size(DIM_1),
124+ q.size(DIM_2) / ori_kv.size(DIM_2)
125+ };
126+ } else {
127+ // 对齐 Python: [ori_kv.shape[1], q.shape[0], q.shape[2] // ori_kv.shape[2]]
128+ if (layout_kv_str == "PA_BBND") {
129+ softmax_lse_size = {
130+ ori_kv.size(DIM_2),
131+ q.size(DIM_0),
132+ q.size(DIM_1) / ori_kv.size(DIM_2)
133+ };
134+ } else {
135+ softmax_lse_size = {
136+ ori_kv.size(DIM_1),
137+ q.size(DIM_0),
138+ q.size(DIM_1) / ori_kv.size(DIM_1)
139+ };
140+ }
141+ }
142+ } else {
143+ // 不返回时tensor传空
144+ softmax_lse_size = {};
145+ }
146+ softmax_lse = at::empty(softmax_lse_size, q.options().dtype(torch::kFloat32));
147+ 
148+ return std::tuple<at::Tensor, at::Tensor>(atten_out, softmax_lse);
149+}
150+ 
151+std::tuple<at::Tensor, at::Tensor> sparse_flash_mla(
152+ const at::Tensor &q,
153+ const c10::optional<at::Tensor> &ori_kv, const c10::optional<at::Tensor> &cmp_kv,
154+ const c10::optional<at::Tensor> &ori_sparse_indices, const c10::optional<at::Tensor> &cmp_sparse_indices,
155+ const c10::optional<at::Tensor> &ori_block_table, const c10::optional<at::Tensor> &cmp_block_table,
156+ const c10::optional<at::Tensor> &cu_seqlens_q, const c10::optional<at::Tensor> &cu_seqlens_ori_kv,
157+ const c10::optional<at::Tensor> &cu_seqlens_cmp_kv, const c10::optional<at::Tensor> &seqused_q,
158+ const c10::optional<at::Tensor> &seqused_ori_kv, const c10::optional<at::Tensor> &seqused_cmp_kv,
159+ const c10::optional<at::Tensor> &cmp_residual_kv,
160+ const c10::optional<at::Tensor> &ori_topk_length, const c10::optional<at::Tensor> &cmp_topk_length,
161+ const c10::optional<at::Tensor> &sinks, const c10::optional<at::Tensor> &metadata,
162+ double softmax_scale, int64_t cmp_ratio,
163+ int64_t ori_mask_mode, int64_t cmp_mask_mode,
164+ int64_t ori_win_left, int64_t ori_win_right,
165+ c10::string_view layout_q, c10::string_view layout_kv,
166+ int64_t topk_value_mode, bool return_softmax_lse)
167+{
168+ TORCH_CHECK(q.numel() > 0, "Tensor query is empty.")
169+ 
170+ std::string layout_q_str = std::string(layout_q);
171+ std::string layout_kv_str = std::string(layout_kv);
172+ const at::Tensor& ori_kv_val = *ori_kv;
173+ // convert str
174+ char *layout_q_ptr = const_cast<char *>(layout_q_str.c_str());
175+ char *layout_kv_ptr = const_cast<char *>(layout_kv_str.c_str());
176+
177+ // construct the atten_out tensor
178+ std::tuple<at::Tensor, at::Tensor> sparse_flash_mla_atten_out = op_api::construct_sparse_flash_mla_atten_out_tensor(
179+ q, ori_kv_val, layout_q_str, layout_kv_str, 64, return_softmax_lse);
180+ at::Tensor atten_out = std::get<0>(sparse_flash_mla_atten_out);
181+ at::Tensor softmax_lse = std::get<1>(sparse_flash_mla_atten_out);
182+ 
183+ ACLNN_CMD(aclnnSparseFlashMla, q,
184+ ori_kv, cmp_kv,
185+ ori_sparse_indices, cmp_sparse_indices,
186+ ori_block_table, cmp_block_table,
187+ cu_seqlens_q, cu_seqlens_ori_kv,
188+ cu_seqlens_cmp_kv, seqused_q,
189+ seqused_ori_kv, seqused_cmp_kv,
190+ cmp_residual_kv,
191+ ori_topk_length, cmp_topk_length,
192+ sinks, metadata,
193+ softmax_scale, cmp_ratio,
194+ ori_mask_mode, cmp_mask_mode,
195+ ori_win_left, ori_win_right,
196+ layout_q_ptr, layout_kv_ptr,
197+ topk_value_mode, return_softmax_lse,
198+ atten_out, softmax_lse);
199+ return std::tuple<at::Tensor, at::Tensor>(atten_out, softmax_lse);
200+}
201+ 
202+PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
203+{
204+ m.def("sparse_flash_mla_metadata", &sparse_flash_mla_metadata, "sparse_flash_mla_metadata");
205+ m.def("sparse_flash_mla", &sparse_flash_mla, "sparse_flash_mla");
206+}
207+} // namespace op_api
Rtorch_extension/npu_ops_transformer/ops/csrc/sparse_flash_mla_grad.cpptorch_extension/cann_ops_transformer/ops/csrc/sparse_flash_mla_grad.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/sparse_flash_mla_grad_metadata.cpptorch_extension/cann_ops_transformer/ops/csrc/sparse_flash_mla_grad_metadata.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/deep_ep.pytorch_extension/cann_ops_transformer/ops/deep_ep.py+4-4
@@ -11,8 +11,8 @@ import torch
11import torch_npu11import torch_npu
12from torch.library import impl12from torch.library import impl
13from torch_npu.utils._error_code import ErrCode, ops_error13from torch_npu.utils._error_code import ErrCode, ops_error
14-from npu_ops_transformer.op_builder.builder import OpBuilder14+from cann_ops_transformer.op_builder.builder import OpBuilder
15-from npu_ops_transformer.op_builder.builder import AS_LIBRARY15+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
16from .moe_distribute_combine import npu_moe_distribute_combine16from .moe_distribute_combine import npu_moe_distribute_combine
17from .moe_distribute_dispatch import npu_moe_distribute_dispatch17from .moe_distribute_dispatch import npu_moe_distribute_dispatch
18from .comm_context import CommContextManager18from .comm_context import CommContextManager
@@ -104,7 +104,7 @@ class MoeDistributeBuffer:
104 const_expert_num=0, elastic_info=None, expert_shard_type=0, shared_expert_num=1,104 const_expert_num=0, elastic_info=None, expert_shard_type=0, shared_expert_num=1,
105 shared_expert_rank_num=0, expert_token_nums_type=1, num_max_dispatch_tokens_per_rank=0):105 shared_expert_rank_num=0, expert_token_nums_type=1, num_max_dispatch_tokens_per_rank=0):
106 (expand_x, dynamic_scales, expand_idx, expert_token_nums, ep_recv_counts, tp_recv_counts, expand_scales) \106 (expand_x, dynamic_scales, expand_idx, expert_token_nums, ep_recv_counts, tp_recv_counts, expand_scales) \
107- = torch.ops.npu_ops_transformer.npu_moe_distribute_dispatch(107+ = torch.ops.cann_ops_transformer.npu_moe_distribute_dispatch(
108 context=self.context,108 context=self.context,
109 x=x,109 x=x,
110 expert_ids=topk_idx,110 expert_ids=topk_idx,
@@ -135,7 +135,7 @@ class MoeDistributeBuffer:
135 const_expert_alpha_2=None, const_expert_v=None, zero_expert_num=0, copy_expert_num=0,135 const_expert_alpha_2=None, const_expert_v=None, zero_expert_num=0, copy_expert_num=0,
136 const_expert_num=0, expert_shared_type=0, shared_expert_num=1, shared_expert_rank_num=0,136 const_expert_num=0, expert_shared_type=0, shared_expert_num=1, shared_expert_rank_num=0,
137 num_max_dispatch_tokens_per_rank=0):137 num_max_dispatch_tokens_per_rank=0):
138- return torch.ops.npu_ops_transformer.npu_moe_distribute_combine(138+ return torch.ops.cann_ops_transformer.npu_moe_distribute_combine(
139 context=self.context,139 context=self.context,
140 expand_x=x,140 expand_x=x,
141 expert_ids=topk_idx,141 expert_ids=topk_idx,
Rtorch_extension/npu_ops_transformer/ops/flash_attn.pytorch_extension/cann_ops_transformer/ops/flash_attn.py+115-115
@@ -1,116 +1,116 @@
1-# -----------------------------------------------------------------------------------------------------------1+# -----------------------------------------------------------------------------------------------------------
2-# Copyright (c) 2026 Huawei Technologies Co., Ltd.2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3-# This program is free software, you can redistribute it and/or modify it under the terms and conditions of3+# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4-# CANN Open Software License Agreement Version 2.0 (the "License").4+# CANN Open Software License Agreement Version 2.0 (the "License").
5-# Please refer to the License for details. You may not use this file except in compliance with the License.5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6-# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7-# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8-# See LICENSE in the root of the software repository for the full text of the License.8+# See LICENSE in the root of the software repository for the full text of the License.
9-# -----------------------------------------------------------------------------------------------------------9+# -----------------------------------------------------------------------------------------------------------
10-import torch10+import torch
11-import torch_npu11+import torch_npu
12-from torch.library import impl12+from torch.library import impl
13-from npu_ops_transformer.op_builder.builder import OpBuilder13+from cann_ops_transformer.op_builder.builder import OpBuilder
14-from npu_ops_transformer.op_builder.builder import AS_LIBRARY14+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
15- 15+ 
16- 16+ 
17-class FlashAttenOpBuilder(OpBuilder):17+class FlashAttenOpBuilder(OpBuilder):
18- def __init__(self):18+ def __init__(self):
19- super(FlashAttenOpBuilder, self).__init__("npu_flash_attn")19+ super(FlashAttenOpBuilder, self).__init__("npu_flash_attn")
20- 20+ 
21- def sources(self):21+ def sources(self):
22- """Path to C++ source code."""22+ """Path to C++ source code."""
23- return ['ops/csrc/flash_attn.cpp']23+ return ['ops/csrc/flash_attn.cpp']
24- 24+ 
25- def schema(self) -> str:25+ def schema(self) -> str:
26- """PyTorch operator signature."""26+ """PyTorch operator signature."""
27- return "npu_flash_attn(Tensor q, Tensor k, Tensor v," \27+ return "npu_flash_attn(Tensor q, Tensor k, Tensor v," \
28- "Tensor?block_table=None, Tensor?cu_seqlens_q=None," \28+ "Tensor?block_table=None, Tensor?cu_seqlens_q=None," \
29- "Tensor?cu_seqlens_kv=None, Tensor?seqused_q=None," \29+ "Tensor?cu_seqlens_kv=None, Tensor?seqused_q=None," \
30- "Tensor?seqused_kv=None, Tensor?sinks=None, Tensor?attn_mask=None, Tensor?metadata=None," \30+ "Tensor?seqused_kv=None, Tensor?sinks=None, Tensor?attn_mask=None, Tensor?metadata=None," \
31- "float softmax_scale=1.0, int mask_mode=0, int win_left=-1, int win_right=-1," \31+ "float softmax_scale=1.0, int mask_mode=0, int win_left=-1, int win_right=-1," \
32- "int max_seqlen_q=-1, int max_seqlen_kv=-1," \32+ "int max_seqlen_q=-1, int max_seqlen_kv=-1," \
33- "str layout_q=\"BSND\", str layout_kv=\"BSND\", str layout_out=\"BSND\"," \33+ "str layout_q=\"BSND\", str layout_kv=\"BSND\", str layout_out=\"BSND\"," \
34- "int return_softmax_lse=0)->(Tensor, Tensor)"34+ "int return_softmax_lse=0)->(Tensor, Tensor)"
35- 35+ 
36- def register_meta(self):36+ def register_meta(self):
37- """37+ """
38- Registers the Meta implementation (Shape/Dtype inference).38+ Registers the Meta implementation (Shape/Dtype inference).
39- Essential for Autograd and FakeTensor support.39+ Essential for Autograd and FakeTensor support.
40- """40+ """
41- @impl(AS_LIBRARY, self.name, "Meta")41+ @impl(AS_LIBRARY, self.name, "Meta")
42- def npu_flash_attn_meta(q, k, v, block_table=None, cu_seqlens_q=None,42+ def npu_flash_attn_meta(q, k, v, block_table=None, cu_seqlens_q=None,
43- cu_seqlens_kv=None, seqused_q=None,43+ cu_seqlens_kv=None, seqused_q=None,
44- seqused_kv=None, sinks=None, attn_mask=None, metadata=None,44+ seqused_kv=None, sinks=None, attn_mask=None, metadata=None,
45- softmax_scale=1.0, mask_mode=0, win_left=-1, win_right=-1,45+ softmax_scale=1.0, mask_mode=0, win_left=-1, win_right=-1,
46- max_seqlen_q=-1, max_seqlen_kv=-1,46+ max_seqlen_q=-1, max_seqlen_kv=-1,
47- layout_q="BSND", layout_kv="BSND", layout_out="BSND",47+ layout_q="BSND", layout_kv="BSND", layout_out="BSND",
48- return_softmax_lse=0):48+ return_softmax_lse=0):
49- if layout_q == "TND":49+ if layout_q == "TND":
50- tSize = q.size(0)50+ tSize = q.size(0)
51- nSize = q.size(1)51+ nSize = q.size(1)
52- dSize = v.size(2)52+ dSize = v.size(2)
53- softmaxOutSize = (nSize, tSize)53+ softmaxOutSize = (nSize, tSize)
54- elif layout_q == "BSND":54+ elif layout_q == "BSND":
55- bSize = q.size(0)55+ bSize = q.size(0)
56- sSize = q.size(1)56+ sSize = q.size(1)
57- nSize = q.size(2)57+ nSize = q.size(2)
58- dSize = v.size(3)58+ dSize = v.size(3)
59- softmaxOutSize = (bSize, nSize, sSize)59+ softmaxOutSize = (bSize, nSize, sSize)
60- else:60+ else:
61- bSize = q.size(0)61+ bSize = q.size(0)
62- nSize = q.size(1)62+ nSize = q.size(1)
63- sSize = q.size(2)63+ sSize = q.size(2)
64- dSize = v.size(3)64+ dSize = v.size(3)
65- softmaxOutSize = (bSize, nSize, sSize)65+ softmaxOutSize = (bSize, nSize, sSize)
66- 66+ 
67- if layout_out == "TND":67+ if layout_out == "TND":
68- torch._check(68+ torch._check(
69- layout_q == "TND",69+ layout_q == "TND",
70- lambda: "When the layout of output is TND, the layout of query must be TND, but got " + str(layout_q),70+ lambda: "When the layout of output is TND, the layout of query must be TND, but got " + str(layout_q),
71- )71+ )
72- attentionOutSize = (tSize, nSize, dSize)72+ attentionOutSize = (tSize, nSize, dSize)
73- elif layout_out == "BNSD":73+ elif layout_out == "BNSD":
74- torch._check(74+ torch._check(
75- layout_q == "BNSD",75+ layout_q == "BNSD",
76- lambda: "When the layout of output is BNSD, the layout of query must be BNSD, but got " + str(layout_q),76+ lambda: "When the layout of output is BNSD, the layout of query must be BNSD, but got " + str(layout_q),
77- )77+ )
78- attentionOutSize = (bSize, nSize, sSize, dSize)78+ attentionOutSize = (bSize, nSize, sSize, dSize)
79- else:79+ else:
80- torch._check(80+ torch._check(
81- layout_q != "TND",81+ layout_q != "TND",
82- lambda: "When the layout of output is BSND, the layout of query must be BNSD or BSND, but got " + str(layout_q),82+ lambda: "When the layout of output is BSND, the layout of query must be BNSD or BSND, but got " + str(layout_q),
83- )83+ )
84- attentionOutSize = (bSize, sSize, nSize, dSize)84+ attentionOutSize = (bSize, sSize, nSize, dSize)
85- 85+ 
86- 86+
87- return (87+ return (
88- torch.empty(attentionOutSize, dtype=q.dtype, device='meta'),88+ torch.empty(attentionOutSize, dtype=q.dtype, device='meta'),
89- torch.empty(softmaxOutSize, dtype=q.dtype, device='meta')89+ torch.empty(softmaxOutSize, dtype=q.dtype, device='meta')
90- )90+ )
91- 91+ 
92- 92+ 
93-# Instantiate the builder93+# Instantiate the builder
94-flash_attn_op_builder = FlashAttenOpBuilder()94+flash_attn_op_builder = FlashAttenOpBuilder()
95-op_module = flash_attn_op_builder.load() # Compiles/loads the .so file95+op_module = flash_attn_op_builder.load() # Compiles/loads the .so file
96- 96+ 
97- 97+ 
98-@impl(AS_LIBRARY, flash_attn_op_builder.name, "PrivateUse1")98+@impl(AS_LIBRARY, flash_attn_op_builder.name, "PrivateUse1")
99-def npu_flash_attn(q, k, v, block_table=None, cu_seqlens_q=None,99+def npu_flash_attn(q, k, v, block_table=None, cu_seqlens_q=None,
100- cu_seqlens_kv=None, seqused_q=None,100+ cu_seqlens_kv=None, seqused_q=None,
101- seqused_kv=None, sinks=None, attn_mask=None, metadata=None,101+ seqused_kv=None, sinks=None, attn_mask=None, metadata=None,
102- softmax_scale=1.0, mask_mode=0, win_left=-1, win_right=-1,102+ softmax_scale=1.0, mask_mode=0, win_left=-1, win_right=-1,
103- max_seqlen_q=-1, max_seqlen_kv=-1,103+ max_seqlen_q=-1, max_seqlen_kv=-1,
104- layout_q="BSND", layout_kv="BSND", layout_out="BSND",104+ layout_q="BSND", layout_kv="BSND", layout_out="BSND",
105- return_softmax_lse=0):105+ return_softmax_lse=0):
106- """106+ """
107- dispatcher implementation for NPU.107+ dispatcher implementation for NPU.
108- 'PrivateUse1' is the combine key for custom NPU backends.108+ 'PrivateUse1' is the combine key for custom NPU backends.
109- """109+ """
110- return op_module.npu_flash_attn(q, k, v, block_table, cu_seqlens_q,110+ return op_module.npu_flash_attn(q, k, v, block_table, cu_seqlens_q,
111- cu_seqlens_kv, seqused_q,111+ cu_seqlens_kv, seqused_q,
112- seqused_kv, sinks, attn_mask, metadata,112+ seqused_kv, sinks, attn_mask, metadata,
113- softmax_scale, mask_mode, win_left, win_right,113+ softmax_scale, mask_mode, win_left, win_right,
114- max_seqlen_q, max_seqlen_kv,114+ max_seqlen_q, max_seqlen_kv,
115- layout_q, layout_kv, layout_out,115+ layout_q, layout_kv, layout_out,
116 return_softmax_lse)116 return_softmax_lse)
Rtorch_extension/npu_ops_transformer/ops/flash_attn_metadata.pytorch_extension/cann_ops_transformer/ops/flash_attn_metadata.py+116-116
@@ -1,116 +1,116 @@
1-# -----------------------------------------------------------------------------------------------------------1+# -----------------------------------------------------------------------------------------------------------
2-# Copyright (c) 2026 Huawei Technologies Co., Ltd.2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3-# This program is free software, you can redistribute it and/or modify it under the terms and conditions of3+# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4-# CANN Open Software License Agreement Version 2.0 (the "License").4+# CANN Open Software License Agreement Version 2.0 (the "License").
5-# Please refer to the License for details. You may not use this file except in compliance with the License.5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6-# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7-# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8-# See LICENSE in the root of the software repository for the full text of the License.8+# See LICENSE in the root of the software repository for the full text of the License.
9-# -----------------------------------------------------------------------------------------------------------9+# -----------------------------------------------------------------------------------------------------------
10-import torch10+import torch
11-import torch_npu11+import torch_npu
12-from torch.library import impl12+from torch.library import impl
13-from npu_ops_transformer.op_builder.builder import OpBuilder13+from cann_ops_transformer.op_builder.builder import OpBuilder
14-from npu_ops_transformer.op_builder.builder import AS_LIBRARY14+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
15-from typing import Optional15+from typing import Optional
16- 16+ 
17-def _calculate_batch_size(batch_size, cu_seqlens_q, seqused_q):17+def _calculate_batch_size(batch_size, cu_seqlens_q, seqused_q):
18- batchSize = 0 18+ batchSize = 0
19- if batch_size is not None: 19+ if batch_size is not None:
20- batchSize = batch_size 20+ batchSize = batch_size
21- elif cu_seqlens_q is not None and cu_seqlens_q.size(0) > 0: 21+ elif cu_seqlens_q is not None and cu_seqlens_q.size(0) > 0:
22- batchSize = cu_seqlens_q.size(0) - 1 22+ batchSize = cu_seqlens_q.size(0) - 1
23- elif seqused_q is not None: 23+ elif seqused_q is not None:
24- batchSize = seqused_q.size(0)24+ batchSize = seqused_q.size(0)
25- return batchSize25+ return batchSize
26- 26+ 
27-def _calculate_metadata_size(batch_size, num_heads_kv): 27+def _calculate_metadata_size(batch_size, num_heads_kv):
28- """计算 metadata tensor 的对齐后大小""" 28+ """计算 metadata tensor 的对齐后大小"""
29- metadataSize = ((36 + 72) * batch_size * num_heads_kv + 1) * 16 29+ metadataSize = ((36 + 72) * batch_size * num_heads_kv + 1) * 16
30- alignedSize = ((metadataSize + 4095) // 4096) * 4096 30+ alignedSize = ((metadataSize + 4095) // 4096) * 4096
31- return alignedSize31+ return alignedSize
32- 32+ 
33-class FlashAttnMetadataOpBuilder(OpBuilder):33+class FlashAttnMetadataOpBuilder(OpBuilder):
34- def __init__(self):34+ def __init__(self):
35- super(FlashAttnMetadataOpBuilder, self).__init__("npu_flash_attn_metadata")35+ super(FlashAttnMetadataOpBuilder, self).__init__("npu_flash_attn_metadata")
36- 36+ 
37- def sources(self):37+ def sources(self):
38- """Path to C++ source code."""38+ """Path to C++ source code."""
39- return ['ops/csrc/flash_attn_metadata.cpp']39+ return ['ops/csrc/flash_attn_metadata.cpp']
40- 40+ 
41- def schema(self) -> str:41+ def schema(self) -> str:
42- """PyTorch operator signature."""42+ """PyTorch operator signature."""
43- return "npu_flash_attn_metadata( int num_heads_q, int num_heads_kv, int head_dim, *, " \43+ return "npu_flash_attn_metadata( int num_heads_q, int num_heads_kv, int head_dim, *, " \
44- "Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_kv=None, Tensor? seqused_q=None, Tensor? seqused_kv=None," \44+ "Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_kv=None, Tensor? seqused_q=None, Tensor? seqused_kv=None," \
45- "int? batch_size=None, int? max_seqlen_q=None, int? max_seqlen_kv=None, " \45+ "int? batch_size=None, int? max_seqlen_q=None, int? max_seqlen_kv=None, " \
46- "int? mask_mode=None, int? win_left=None, int? win_right=None, " \46+ "int? mask_mode=None, int? win_left=None, int? win_right=None, " \
47- "str? layout_q=None, str? layout_kv=None, str? layout_out=None) -> Tensor"47+ "str? layout_q=None, str? layout_kv=None, str? layout_out=None) -> Tensor"
48- 48+ 
49- def register_meta(self):49+ def register_meta(self):
50- """50+ """
51- Registers Meta implementation (Shape/Dtype inference).51+ Registers Meta implementation (Shape/Dtype inference).
52- Essential for Autograd and FakeTensor support.52+ Essential for Autograd and FakeTensor support.
53- """53+ """
54- @torch.library.register_fake("npu_ops_transformer::" + self.name)54+ @torch.library.register_fake("cann_ops_transformer::" + self.name)
55- def npu_flash_attn_metadata_meta( num_heads_q, num_heads_kv, head_dim,55+ def npu_flash_attn_metadata_meta( num_heads_q, num_heads_kv, head_dim,
56- cu_seqlens_q = None, cu_seqlens_kv = None, seqused_q = None, seqused_kv = None,56+ cu_seqlens_q = None, cu_seqlens_kv = None, seqused_q = None, seqused_kv = None,
57- batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None,57+ batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None,
58- max_seqlen_kv: Optional[int] = None,58+ max_seqlen_kv: Optional[int] = None,
59- mask_mode: Optional[int] = None, win_left: Optional[int] = None,59+ mask_mode: Optional[int] = None, win_left: Optional[int] = None,
60- win_right: Optional[int] = None, layout_q: Optional[str] = None,60+ win_right: Optional[int] = None, layout_q: Optional[str] = None,
61- layout_kv: Optional[str] = None, layout_out: Optional[str] = None):61+ layout_kv: Optional[str] = None, layout_out: Optional[str] = None):
62- metadataSize = _calculate_metadata_size(batch_size, num_heads_kv)62+ metadataSize = _calculate_metadata_size(batch_size, num_heads_kv)
63- return torch.empty((metadataSize,), dtype=torch.int32, device="npu")63+ return torch.empty((metadataSize,), dtype=torch.int32, device="npu")
64- 64+ 
65-# Instantiate the builder65+# Instantiate the builder
66-flash_attn_metadata_op_builder = FlashAttnMetadataOpBuilder()66+flash_attn_metadata_op_builder = FlashAttnMetadataOpBuilder()
67-op_module = flash_attn_metadata_op_builder.load()67+op_module = flash_attn_metadata_op_builder.load()
68- 68+ 
69- 69+ 
70-@impl(AS_LIBRARY, flash_attn_metadata_op_builder.name, "PrivateUse1")70+@impl(AS_LIBRARY, flash_attn_metadata_op_builder.name, "PrivateUse1")
71-def npu_flash_attn_metadata(num_heads_q, num_heads_kv, head_dim,71+def npu_flash_attn_metadata(num_heads_q, num_heads_kv, head_dim,
72- cu_seqlens_q = None, cu_seqlens_kv = None, seqused_q = None, seqused_kv = None,72+ cu_seqlens_q = None, cu_seqlens_kv = None, seqused_q = None, seqused_kv = None,
73- batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None,73+ batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None,
74- max_seqlen_kv: Optional[int] = None,74+ max_seqlen_kv: Optional[int] = None,
75- mask_mode: Optional[int] = None, win_left: Optional[int] = None,75+ mask_mode: Optional[int] = None, win_left: Optional[int] = None,
76- win_right: Optional[int] = None, layout_q: Optional[str] = None,76+ win_right: Optional[int] = None, layout_q: Optional[str] = None,
77- layout_kv: Optional[str] = None, layout_out: Optional[str] = None):77+ layout_kv: Optional[str] = None, layout_out: Optional[str] = None):
78- """78+ """
79- Dispatcher implementation: NPU.79+ Dispatcher implementation: NPU.
80- 'PrivateUse1' is dispatch key for custom NPU backends.80+ 'PrivateUse1' is dispatch key for custom NPU backends.
81- """81+ """
82- batch_size = _calculate_batch_size(batch_size, cu_seqlens_q, seqused_q) if batch_size is None else batch_size82+ batch_size = _calculate_batch_size(batch_size, cu_seqlens_q, seqused_q) if batch_size is None else batch_size
83- max_seqlen_q = -1 if max_seqlen_q is None else max_seqlen_q83+ max_seqlen_q = -1 if max_seqlen_q is None else max_seqlen_q
84- max_seqlen_kv = -1 if max_seqlen_kv is None else max_seqlen_kv84+ max_seqlen_kv = -1 if max_seqlen_kv is None else max_seqlen_kv
85- mask_mode = 1 if mask_mode is None else mask_mode85+ mask_mode = 1 if mask_mode is None else mask_mode
86- win_left = -1 if win_left is None else win_left86+ win_left = -1 if win_left is None else win_left
87- win_right = -1 if win_right is None else win_right87+ win_right = -1 if win_right is None else win_right
88- layout_q = "BSND" if layout_q is None else layout_q88+ layout_q = "BSND" if layout_q is None else layout_q
89- layout_kv = "BSND" if layout_kv is None else layout_kv89+ layout_kv = "BSND" if layout_kv is None else layout_kv
90- layout_out = "BSND" if layout_out is None else layout_out90+ layout_out = "BSND" if layout_out is None else layout_out
91- 91+ 
92- metadataSize = _calculate_metadata_size(batch_size, num_heads_kv)92+ metadataSize = _calculate_metadata_size(batch_size, num_heads_kv)
93- output = torch.empty((metadataSize,), dtype=torch.int32, device="npu")93+ output = torch.empty((metadataSize,), dtype=torch.int32, device="npu")
94- 94+
95- return op_module.npu_flash_attn_metadata(cu_seqlens_q, cu_seqlens_kv, seqused_q, seqused_kv,95+ return op_module.npu_flash_attn_metadata(cu_seqlens_q, cu_seqlens_kv, seqused_q, seqused_kv,
96- num_heads_q, num_heads_kv, head_dim,96+ num_heads_q, num_heads_kv, head_dim,
97- batch_size, max_seqlen_q, max_seqlen_kv,97+ batch_size, max_seqlen_q, max_seqlen_kv,
98- mask_mode, win_left, win_right, layout_q, layout_kv, layout_out, output)98+ mask_mode, win_left, win_right, layout_q, layout_kv, layout_out, output)
99- 99+ 
100- 100+ 
101-@torch.library.register_kernel("npu_ops_transformer::npu_flash_attn_metadata", None) 101+@torch.library.register_kernel("cann_ops_transformer::npu_flash_attn_metadata", None)
102-def npu_flash_attn_metadata_fallback(num_heads_q, num_heads_kv, head_dim,102+def npu_flash_attn_metadata_fallback(num_heads_q, num_heads_kv, head_dim,
103- cu_seqlens_q = None, cu_seqlens_kv = None, seqused_q = None, seqused_kv = None,103+ cu_seqlens_q = None, cu_seqlens_kv = None, seqused_q = None, seqused_kv = None,
104- batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None,104+ batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None,
105- max_seqlen_kv: Optional[int] = None,105+ max_seqlen_kv: Optional[int] = None,
106- mask_mode: Optional[int] = None, win_left: Optional[int] = None,106+ mask_mode: Optional[int] = None, win_left: Optional[int] = None,
107- win_right: Optional[int] = None, layout_q: Optional[str] = None,107+ win_right: Optional[int] = None, layout_q: Optional[str] = None,
108- layout_kv: Optional[str] = None, layout_out: Optional[str] = None): 108+ layout_kv: Optional[str] = None, layout_out: Optional[str] = None):
109- # 处理所有 tensor 都为 None 的情况 109+ # 处理所有 tensor 都为 None 的情况
110- return npu_flash_attn_metadata(num_heads_q, num_heads_kv, head_dim,110+ return npu_flash_attn_metadata(num_heads_q, num_heads_kv, head_dim,
111- cu_seqlens_q, cu_seqlens_kv, seqused_q, seqused_kv,111+ cu_seqlens_q, cu_seqlens_kv, seqused_q, seqused_kv,
112- batch_size, max_seqlen_q,112+ batch_size, max_seqlen_q,
113- max_seqlen_kv,113+ max_seqlen_kv,
114- mask_mode, win_left,114+ mask_mode, win_left,
115- win_right, layout_q,115+ win_right, layout_q,
116- layout_kv, layout_out) 116+ layout_kv, layout_out)
Rtorch_extension/npu_ops_transformer/ops/graph_convert/graph_convert_flash_attn.pytorch_extension/cann_ops_transformer/ops/graph_convert/graph_convert_flash_attn.py+101-101
@@ -1,102 +1,102 @@
1-try:1+try:
2- import torch2+ import torch
3- import torch_npu3+ import torch_npu
4- import torchair4+ import torchair
5- from torch.library import impl5+ from torch.library import impl
6- from torchair._ge_concrete_graph import ge_apis as ge6+ from torchair._ge_concrete_graph import ge_apis as ge
7- from torchair.ge._ge_graph import Tensor, TensorSpec7+ from torchair.ge._ge_graph import Tensor, TensorSpec
8- from torchair._ge_concrete_graph.fx2ge_converter import declare_supported, register_fx_node_ge_converter8+ from torchair._ge_concrete_graph.fx2ge_converter import declare_supported, register_fx_node_ge_converter
9- from torchair._ge_concrete_graph.supported_declaration import Support9+ from torchair._ge_concrete_graph.supported_declaration import Support
10- from typing import Any, Dict, List, Tuple, Union, Callable, Optional10+ from typing import Any, Dict, List, Tuple, Union, Callable, Optional
11- from torchair._ge_concrete_graph.ge_ir_pb2 import GraphDef, OpDef, TensorDescriptor, TensorDef11+ from torchair._ge_concrete_graph.ge_ir_pb2 import GraphDef, OpDef, TensorDescriptor, TensorDef
12- from torchair.ge._ge_graph import get_default_ge_graph, next_unique_name12+ from torchair.ge._ge_graph import get_default_ge_graph, next_unique_name
13- from torchair.ge._ge_graph import auto_convert_to_tensor13+ from torchair.ge._ge_graph import auto_convert_to_tensor
14- from torchair.ge._ge_graph import Tensor, TensorSpec, DataType, TensorType14+ from torchair.ge._ge_graph import Tensor, TensorSpec, DataType, TensorType
15- from torchair.ge._ge_graph import compat_as_bytes, compat_as_bytes_list15+ from torchair.ge._ge_graph import compat_as_bytes, compat_as_bytes_list
16- from torchair.ge._ge_graph import trans_to_list_list_int, trans_to_list_list_float16+ from torchair.ge._ge_graph import trans_to_list_list_int, trans_to_list_list_float
17- from torchair.ge._ge_graph import get_invalid_desc17+ from torchair.ge._ge_graph import get_invalid_desc
18- from torchair._ge_concrete_graph.compat_ir import ge_op, IrDef18+ from torchair._ge_concrete_graph.compat_ir import ge_op, IrDef
19- from torchair.ge import attr19+ from torchair.ge import attr
20- _TORCHAIR_AVAILABLE = True20+ _TORCHAIR_AVAILABLE = True
21-except ImportError:21+except ImportError:
22- _TORCHAIR_AVAILABLE = False22+ _TORCHAIR_AVAILABLE = False
23- 23+ 
24-if _TORCHAIR_AVAILABLE:24+if _TORCHAIR_AVAILABLE:
25- @auto_convert_to_tensor(25+ @auto_convert_to_tensor(
26- [False, False, False, False, False, False, False, False, False, False, False],26+ [False, False, False, False, False, False, False, False, False, False, False],
27- [False, False, False, True, True, True, True, True, True, True])27+ [False, False, False, True, True, True, True, True, True, True])
28- def FlashAttn(q: Tensor,28+ def FlashAttn(q: Tensor,
29- k: Tensor,29+ k: Tensor,
30- v: Tensor,30+ v: Tensor,
31- block_table: Optional[Tensor],31+ block_table: Optional[Tensor],
32- cu_seqlens_q: Optional[Tensor],32+ cu_seqlens_q: Optional[Tensor],
33- cu_seqlens_kv: Optional[Tensor],33+ cu_seqlens_kv: Optional[Tensor],
34- seqused_q: Optional[Tensor],34+ seqused_q: Optional[Tensor],
35- seqused_kv: Optional[Tensor],35+ seqused_kv: Optional[Tensor],
36- sinks: Optional[Tensor],36+ sinks: Optional[Tensor],
37- attn_mask: Optional[Tensor],37+ attn_mask: Optional[Tensor],
38- metadata: Optional[Tensor],38+ metadata: Optional[Tensor],
39- softmax_scale: float = 1.0,39+ softmax_scale: float = 1.0,
40- mask_mode: int = 0,40+ mask_mode: int = 0,
41- win_left: int = -1,41+ win_left: int = -1,
42- win_right: int = -1,42+ win_right: int = -1,
43- max_seqlen_q: int = -1,43+ max_seqlen_q: int = -1,
44- max_seqlen_kv: int = -1,44+ max_seqlen_kv: int = -1,
45- layout_q: str = "BSND",45+ layout_q: str = "BSND",
46- layout_kv: str = "BSND",46+ layout_kv: str = "BSND",
47- layout_out: str = "BSND",47+ layout_out: str = "BSND",
48- return_softmax_lse: int = 0,48+ return_softmax_lse: int = 0,
49- deterministic: int = 0):49+ deterministic: int = 0):
50- 50+
51- result = q.new_empty(q.size())51+ result = q.new_empty(q.size())
52- return result52+ return result
53- 53+
54- @register_fx_node_ge_converter(torch.ops.npu_ops_transformer.npu_flash_attn.default)54+ @register_fx_node_ge_converter(torch.ops.cann_ops_transformer.npu_flash_attn.default)
55- def convert_npu_flash_attn(55+ def convert_npu_flash_attn(
56- q: Tensor,56+ q: Tensor,
57- k: Tensor,57+ k: Tensor,
58- v: Tensor,58+ v: Tensor,
59- block_table: Tensor = None,59+ block_table: Tensor = None,
60- cu_seqlens_q: Tensor = None,60+ cu_seqlens_q: Tensor = None,
61- cu_seqlens_kv: Tensor = None,61+ cu_seqlens_kv: Tensor = None,
62- seqused_q: Tensor = None,62+ seqused_q: Tensor = None,
63- seqused_kv: Tensor = None,63+ seqused_kv: Tensor = None,
64- sinks: Tensor = None,64+ sinks: Tensor = None,
65- attn_mask: Tensor = None,65+ attn_mask: Tensor = None,
66- metadata: Tensor = None,66+ metadata: Tensor = None,
67- softmax_scale: float = 1.0,67+ softmax_scale: float = 1.0,
68- mask_mode: int = 0,68+ mask_mode: int = 0,
69- win_left: int = -1,69+ win_left: int = -1,
70- win_right: int = -1,70+ win_right: int = -1,
71- max_seqlen_q: int = -1,71+ max_seqlen_q: int = -1,
72- max_seqlen_kv: int = -1,72+ max_seqlen_kv: int = -1,
73- layout_q: str = "BSND",73+ layout_q: str = "BSND",
74- layout_kv: str = "BSND",74+ layout_kv: str = "BSND",
75- layout_out: str = "BSND",75+ layout_out: str = "BSND",
76- return_softmax_lse: int = 0,76+ return_softmax_lse: int = 0,
77- deterministic: int = 0):77+ deterministic: int = 0):
78- 78+ 
79- 79+
80- raise AssertionError(f"GE not supported!")80+ raise AssertionError(f"GE not supported!")
81- return FlashAttn(q = q,81+ return FlashAttn(q = q,
82- k = k,82+ k = k,
83- v = v,83+ v = v,
84- block_table = block_table,84+ block_table = block_table,
85- cu_seqlens_q = cu_seqlens_q,85+ cu_seqlens_q = cu_seqlens_q,
86- cu_seqlens_kv = cu_seqlens_kv,86+ cu_seqlens_kv = cu_seqlens_kv,
87- seqused_q = seqused_q,87+ seqused_q = seqused_q,
88- seqused_kv = seqused_kv,88+ seqused_kv = seqused_kv,
89- sinks =sinks,89+ sinks =sinks,
90- attn_mask = attn_mask,90+ attn_mask = attn_mask,
91- metadata = metadata,91+ metadata = metadata,
92- softmax_scale = softmax_scale,92+ softmax_scale = softmax_scale,
93- mask_mode = mask_mode,93+ mask_mode = mask_mode,
94- win_left = win_left,94+ win_left = win_left,
95- win_right = win_right,95+ win_right = win_right,
96- max_seqlen_q = max_seqlen_q,96+ max_seqlen_q = max_seqlen_q,
97- max_seqlen_kv = max_seqlen_kv,97+ max_seqlen_kv = max_seqlen_kv,
98- layout_q = layout_q,98+ layout_q = layout_q,
99- layout_kv = layout_kv,99+ layout_kv = layout_kv,
100- layout_out = layout_out,100+ layout_out = layout_out,
101- return_softmax_lse = return_softmax_lse,101+ return_softmax_lse = return_softmax_lse,
102 deterministic = deterministic)102 deterministic = deterministic)
Rtorch_extension/npu_ops_transformer/ops/graph_convert/graph_convert_lightning_indexer_v2_metadata.pytorch_extension/cann_ops_transformer/ops/graph_convert/graph_convert_lightning_indexer.py+3-3
@@ -33,8 +33,8 @@ except ImportError:
33 _TORCHAIR_AVAILABLE = False33 _TORCHAIR_AVAILABLE = False
34 34 
35if _TORCHAIR_AVAILABLE:35if _TORCHAIR_AVAILABLE:
36- @register_fx_node_ge_converter(torch.ops.npu_ops_transformer.npu_lightning_indexer_v2_metadata.default)36+ @register_fx_node_ge_converter(torch.ops.cann_ops_transformer.lightning_indexer_metadata.default)
37- def convert_npu_lightning_indexer_v2_metadata(37+ def convert_lightning_indexer_metadata(
38 num_heads_q: int, num_heads_kv: int, head_dim: int, topk: int, *,38 num_heads_q: int, num_heads_kv: int, head_dim: int, topk: int, *,
39 cu_seqlens_q: Optional[Tensor] = None, cu_seqlens_k: Optional[Tensor] = None,39 cu_seqlens_q: Optional[Tensor] = None, cu_seqlens_k: Optional[Tensor] = None,
40 seqused_q: Optional[Tensor] = None, seqused_k: Optional[Tensor] = None, cmp_residual_k: Optional[Tensor] = None,40 seqused_q: Optional[Tensor] = None, seqused_k: Optional[Tensor] = None, cmp_residual_k: Optional[Tensor] = None,
@@ -42,4 +42,4 @@ if _TORCHAIR_AVAILABLE:
42 layout_q: Optional[str] = None, layout_k: Optional[str] = None, mask_mode: Optional[int] = None,42 layout_q: Optional[str] = None, layout_k: Optional[str] = None, mask_mode: Optional[int] = None,
43 cmp_ratio: Optional[int] = None,43 cmp_ratio: Optional[int] = None,
44 meta_outputs: TensorSpec = None):44 meta_outputs: TensorSpec = None):
45- raise RuntimeError("GE converter doesn't support op: 'npu_lightning_indexer_v2_metadata'")45+ raise RuntimeError("GE converter doesn't support op: 'lightning_indexer_metadata'")
Rtorch_extension/npu_ops_transformer/ops/graph_convert/graph_convert_mega_moe.pytorch_extension/cann_ops_transformer/ops/graph_convert/graph_convert_mega_moe.py+1-1
@@ -147,7 +147,7 @@ if _TORCHAIR_AVAILABLE:
147 .output("expert_token_nums", "DT_INT32")147 .output("expert_token_nums", "DT_INT32")
148 )148 )
149 149 
150- @register_fx_node_ge_converter(torch.ops.npu_ops_transformer.npu_mega_moe.default)150+ @register_fx_node_ge_converter(torch.ops.cann_ops_transformer.npu_mega_moe.default)
151 def convert_npu_mega_moe(151 def convert_npu_mega_moe(
152 context: Tensor,152 context: Tensor,
153 x: Tensor,153 x: Tensor,
@@ -0,0 +1,49 @@
1+# -----------------------------------------------------------------------------------------------------------
2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+# CANN Open Software License Agreement Version 2.0 (the "License").
5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+# See LICENSE in the root of the software repository for the full text of the License.
9+# -----------------------------------------------------------------------------------------------------------
10+# GE Converter for Graph Mode
11+ 
12+try:
13+ import torch
14+ import torch_npu
15+ import torchair
16+ from torch.library import impl
17+ from torchair._ge_concrete_graph import ge_apis as ge
18+ from torchair.ge._ge_graph import Tensor, TensorSpec
19+ from torchair._ge_concrete_graph.fx2ge_converter import declare_supported, register_fx_node_ge_converter
20+ from torchair._ge_concrete_graph.supported_declaration import Support
21+ from typing import Any, Dict, List, Tuple, Union, Callable, Optional
22+ from torchair._ge_concrete_graph.ge_ir_pb2 import GraphDef, OpDef, TensorDescriptor, TensorDef
23+ from torchair.ge._ge_graph import get_default_ge_graph, next_unique_name
24+ from torchair.ge._ge_graph import auto_convert_to_tensor
25+ from torchair.ge._ge_graph import Tensor, TensorSpec, DataType, TensorType
26+ from torchair.ge._ge_graph import compat_as_bytes, compat_as_bytes_list
27+ from torchair.ge._ge_graph import trans_to_list_list_int, trans_to_list_list_float
28+ from torchair.ge._ge_graph import get_invalid_desc
29+ from torchair._ge_concrete_graph.compat_ir import ge_op, IrDef
30+ from torchair.ge import attr
31+ _TORCHAIR_AVAILABLE = True
32+except ImportError:
33+ _TORCHAIR_AVAILABLE = False
34+ 
35+if _TORCHAIR_AVAILABLE:
36+ @register_fx_node_ge_converter(torch.ops.cann_ops_transformer.mixed_quant_sparse_flash_mla_metadata.default)
37+ def convert_mixed_quant_sparse_flash_mla_metadata(
38+ num_heads_q: int, num_heads_kv: int, head_dim: int, quant_mode: int, *, cu_seqlens_q: Optional[Tensor] = None,
39+ cu_seqlens_ori_kv: Optional[Tensor] = None, cu_seqlens_cmp_kv: Optional[Tensor] = None,
40+ seqused_q: Optional[Tensor] = None, seqused_ori_kv: Optional[Tensor] = None,
41+ seqused_cmp_kv: Optional[Tensor] = None, cmp_residual_kv: Optional[Tensor] = None,
42+ ori_topk_length: Optional[Tensor] = None, cmp_topk_length: Optional[Tensor] = None,
43+ batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_ori_kv: Optional[int] = None,
44+ max_seqlen_cmp_kv: Optional[int] = None, ori_topk: Optional[int] = None, cmp_topk: Optional[int] = None,
45+ rope_head_dim: Optional[int] = None, cmp_ratio: Optional[int] = None, ori_mask_mode: Optional[int] = None,
46+ cmp_mask_mode: Optional[int] = None, ori_win_left: Optional[int] = None, ori_win_right: Optional[int] = None,
47+ layout_q: Optional[str] = None, layout_kv: Optional[str] = None, has_ori_kv: Optional[bool] = None,
48+ has_cmp_kv: Optional[bool] = None, meta_outputs: TensorSpec = None):
49+ raise RuntimeError("GE converter doesn't support op: 'mixed_quant_sparse_flash_mla_metadata'")
Rtorch_extension/npu_ops_transformer/ops/graph_convert/graph_convert_moe_distribute_combine.pytorch_extension/cann_ops_transformer/ops/graph_convert/graph_convert_moe_distribute_combine.py+1-1
@@ -253,7 +253,7 @@ if _TORCHAIR_AVAILABLE:
253 253 
254 return x254 return x
255 255 
256- @register_fx_node_ge_converter(torch.ops.npu_ops_transformer.npu_moe_distribute_combine.default)256+ @register_fx_node_ge_converter(torch.ops.cann_ops_transformer.npu_moe_distribute_combine.default)
257 def convert_npu_moe_distribute_combine(257 def convert_npu_moe_distribute_combine(
258 context: Tensor,258 context: Tensor,
259 expand_x: Tensor,259 expand_x: Tensor,
Rtorch_extension/npu_ops_transformer/ops/graph_convert/graph_convert_moe_distribute_dispatch.pytorch_extension/cann_ops_transformer/ops/graph_convert/graph_convert_moe_distribute_dispatch.py+1-1
@@ -162,7 +162,7 @@ if _TORCHAIR_AVAILABLE:
162 .output("expand_scales", "DT_FLOAT")162 .output("expand_scales", "DT_FLOAT")
163 )163 )
164 164 
165- @register_fx_node_ge_converter(torch.ops.npu_ops_transformer.npu_moe_distribute_dispatch.default)165+ @register_fx_node_ge_converter(torch.ops.cann_ops_transformer.npu_moe_distribute_dispatch.default)
166 def converter_moe_distribute_dispatch(166 def converter_moe_distribute_dispatch(
167 context: Tensor,167 context: Tensor,
168 x: Tensor,168 x: Tensor,
Rtorch_extension/npu_ops_transformer/ops/graph_convert/graph_convert_quant_lightning_indexer_v2_metadata.pytorch_extension/cann_ops_transformer/ops/graph_convert/graph_convert_quant_lightning_indexer.py+4-4
@@ -33,13 +33,13 @@ except ImportError:
33 _TORCHAIR_AVAILABLE = False33 _TORCHAIR_AVAILABLE = False
34 34 
35if _TORCHAIR_AVAILABLE:35if _TORCHAIR_AVAILABLE:
36- @register_fx_node_ge_converter(torch.ops.npu_ops_transformer.npu_quant_lightning_indexer_v2_metadata.default)36+ @register_fx_node_ge_converter(torch.ops.cann_ops_transformer.quant_lightning_indexer_metadata.default)
37- def convert_npu_quant_lightning_indexer_v2_metadata(37+ def convert_quant_lightning_indexer_metadata(
38- num_heads_q: int, num_heads_kv: int, head_dim: int, topk: int, q_quant_mode: int, k_quant_mode: int, *,38+ num_heads_q: int, num_heads_kv: int, head_dim: int, topk: int, quant_mode: int, *,
39 cu_seqlens_q: Optional[Tensor] = None, cu_seqlens_k: Optional[Tensor] = None,39 cu_seqlens_q: Optional[Tensor] = None, cu_seqlens_k: Optional[Tensor] = None,
40 seqused_q: Optional[Tensor] = None, seqused_k: Optional[Tensor] = None, cmp_residual_k: Optional[Tensor] = None,40 seqused_q: Optional[Tensor] = None, seqused_k: Optional[Tensor] = None, cmp_residual_k: Optional[Tensor] = None,
41 batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_k: Optional[int] = None,41 batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_k: Optional[int] = None,
42 layout_q: Optional[str] = None, layout_k: Optional[str] = None, mask_mode: Optional[int] = None,42 layout_q: Optional[str] = None, layout_k: Optional[str] = None, mask_mode: Optional[int] = None,
43 cmp_ratio: Optional[int] = None,43 cmp_ratio: Optional[int] = None,
44 meta_outputs: TensorSpec = None):44 meta_outputs: TensorSpec = None):
45- raise RuntimeError("GE converter doesn't support op: 'npu_quant_lightning_indexer_v2_metadata'")45+ raise RuntimeError("GE converter doesn't support op: 'quant_lightning_indexer_metadata'")
@@ -0,0 +1,49 @@
1+# -----------------------------------------------------------------------------------------------------------
2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+# CANN Open Software License Agreement Version 2.0 (the "License").
5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+# See LICENSE in the root of the software repository for the full text of the License.
9+# -----------------------------------------------------------------------------------------------------------
10+# GE Converter for Graph Mode
11+ 
12+try:
13+ import torch
14+ import torch_npu
15+ import torchair
16+ from torch.library import impl
17+ from torchair._ge_concrete_graph import ge_apis as ge
18+ from torchair.ge._ge_graph import Tensor, TensorSpec
19+ from torchair._ge_concrete_graph.fx2ge_converter import declare_supported, register_fx_node_ge_converter
20+ from torchair._ge_concrete_graph.supported_declaration import Support
21+ from typing import Any, Dict, List, Tuple, Union, Callable, Optional
22+ from torchair._ge_concrete_graph.ge_ir_pb2 import GraphDef, OpDef, TensorDescriptor, TensorDef
23+ from torchair.ge._ge_graph import get_default_ge_graph, next_unique_name
24+ from torchair.ge._ge_graph import auto_convert_to_tensor
25+ from torchair.ge._ge_graph import Tensor, TensorSpec, DataType, TensorType
26+ from torchair.ge._ge_graph import compat_as_bytes, compat_as_bytes_list
27+ from torchair.ge._ge_graph import trans_to_list_list_int, trans_to_list_list_float
28+ from torchair.ge._ge_graph import get_invalid_desc
29+ from torchair._ge_concrete_graph.compat_ir import ge_op, IrDef
30+ from torchair.ge import attr
31+ _TORCHAIR_AVAILABLE = True
32+except ImportError:
33+ _TORCHAIR_AVAILABLE = False
34+ 
35+if _TORCHAIR_AVAILABLE:
36+ @register_fx_node_ge_converter(torch.ops.cann_ops_transformer.sparse_flash_mla_metadata.default)
37+ def convert_sparse_flash_mla_metadata(
38+ num_heads_q: int, num_heads_kv: int, head_dim: int, *, cu_seqlens_q: Optional[Tensor] = None,
39+ cu_seqlens_ori_kv: Optional[Tensor] = None, cu_seqlens_cmp_kv: Optional[Tensor] = None,
40+ seqused_q: Optional[Tensor] = None, seqused_ori_kv: Optional[Tensor] = None,
41+ seqused_cmp_kv: Optional[Tensor] = None, cmp_residual_kv: Optional[Tensor] = None,
42+ ori_topk_length: Optional[Tensor] = None, cmp_topk_length: Optional[Tensor] = None,
43+ batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_ori_kv: Optional[int] = None,
44+ max_seqlen_cmp_kv: Optional[int] = None, ori_topk: Optional[int] = None, cmp_topk: Optional[int] = None,
45+ cmp_ratio: Optional[int] = None, ori_mask_mode: Optional[int] = None, cmp_mask_mode: Optional[int] = None,
46+ ori_win_left: Optional[int] = None, ori_win_right: Optional[int] = None, layout_q: Optional[str] = None,
47+ layout_kv: Optional[str] = None, has_ori_kv: Optional[bool] = None, has_cmp_kv: Optional[bool] = None,
48+ meta_outputs: TensorSpec = None):
49+ raise RuntimeError("GE converter doesn't support op: 'sparse_flash_mla_metadata'")
Rtorch_extension/npu_ops_transformer/ops/lightning_indexer_v2.pytorch_extension/cann_ops_transformer/ops/lightning_indexer_v2.py+144-81
@@ -1,82 +1,145 @@
1-# -----------------------------------------------------------------------------------------------------------1+# -----------------------------------------------------------------------------------------------------------
2-# Copyright (c) 2026 Huawei Technologies Co., Ltd.2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3-# This program is free software, you can redistribute it and/or modify it under the terms and conditions of3+# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4-# CANN Open Software License Agreement Version 2.0 (the "License").4+# CANN Open Software License Agreement Version 2.0 (the "License").
5-# Please refer to the License for details. You may not use this file except in compliance with the License.5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6-# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7-# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8-# See LICENSE in the root of the software repository for the full text of the License.8+# See LICENSE in the root of the software repository for the full text of the License.
9-# -----------------------------------------------------------------------------------------------------------9+# -----------------------------------------------------------------------------------------------------------
10-import torch10+from typing import Optional
11-import torch_npu11+import torch
12-from torch.library import impl12+import torch_npu
13-from npu_ops_transformer.op_builder.builder import OpBuilder13+from torch.library import impl
14-from npu_ops_transformer.op_builder.builder import AS_LIBRARY14+from cann_ops_transformer.op_builder.builder import OpBuilder
15- 15+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
16- 16+LI_V2_METADATA_SIZE = 1024
17-class LightningIndexerV2OpBuilder(OpBuilder):17+LI_V2_METADATA_OP_NAME = "lightning_indexer_metadata"
18- def __init__(self):18+ 
19- super(LightningIndexerV2OpBuilder, self).__init__("npu_lightning_indexer_v2")19+ 
20- 20+class LightningIndexerV2OpBuilder(OpBuilder):
21- def sources(self):21+ def __init__(self):
22- """Path to C++ source code."""22+ super(LightningIndexerV2OpBuilder, self).__init__("lightning_indexer_v2")
23- return ['ops/csrc/lightning_indexer_v2.cpp']23+ 
24- 24+ def sources(self):
25- def schema(self) -> str:25+ """Path to C++ source code."""
26- """PyTorch operator signature."""26+ return ['ops/csrc/lightning_indexer_v2.cpp']
27- return "npu_lightning_indexer_v2(Tensor q, Tensor k, Tensor w, " \27+ 
28- "int topk, *, Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_k=None," \28+ def schema(self) -> str:
29- "Tensor? seqused_q=None, Tensor? seqused_k=None, " \29+ """PyTorch operator signature."""
30- "Tensor? cmp_residual_k=None, Tensor? block_table=None, " \30+ return [
31- "Tensor? output_idx_offset=None, Tensor? metadata=None, int max_seqlen_q=-1," \31+ "lightning_indexer_metadata(int num_heads_q, int num_heads_k, int head_dim, int topk, *, "
32- "str layout_q=\"BSND\", str layout_k=\"BSND\", int mask_mode=0, " \32+ "Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_k=None, Tensor? seqused_q=None, Tensor? seqused_k=None, "
33- "int cmp_ratio=1," \33+ "Tensor? cmp_residual_k=None, int? batch_size=None, int? max_seqlen_q=None, int? max_seqlen_k=None, "
34- "int return_value=0) -> (Tensor, Tensor)"34+ "str? layout_q=None, str? layout_k=None, int? mask_mode=None, int? cmp_ratio=None) -> Tensor",
35- 35+ 
36- def register_meta(self):36+ "lightning_indexer_v2(Tensor q, Tensor k, Tensor w, "
37- """37+ "int topk, *, Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_k=None,"
38- Registers the Meta implementation (Shape/Dtype inference).38+ "Tensor? seqused_q=None, Tensor? seqused_k=None, "
39- Essential for Autograd and FakeTensor support.39+ "Tensor? cmp_residual_k=None, Tensor? block_table=None, "
40- """40+ "Tensor? output_idx_offset=None, Tensor? metadata=None, int max_seqlen_q=-1,"
41- @impl(AS_LIBRARY, self.name, "Meta")41+ "str layout_q=\"BSND\", str layout_k=\"BSND\", int mask_mode=0, "
42- def npu_lightning_indexer_v2_meta(q, k, w, topk, *, cu_seqlens_q=None, cu_seqlens_k=None,42+ "int cmp_ratio=1,"
43- seqused_q=None, seqused_k=None, cmp_residual_k=None, block_table=None,43+ "int return_value=0) -> (Tensor, Tensor)"
44- output_idx_offset=None, metadata=None, max_seqlen_q=-1,44+ ]
45- layout_q="BSND", layout_k="BSND", mask_mode=0,45+ 
46- cmp_ratio=1, return_value=0):46+ def register_meta(self):
47- key_head_num = k.shape[1] if layout_k == "TND" else k.shape[2]47+ """
48- 48+ Registers the Meta implementation (Shape/Dtype inference).
49- if layout_q == "BSND":49+ Essential for Autograd and FakeTensor support.
50- sparse_indices_out = torch.empty([q.shape[0], q.shape[1], key_head_num, topk],50+ """
51- dtype=torch.int32, device="meta")51+ @torch.library.register_fake("cann_ops_transformer::" + LI_V2_METADATA_OP_NAME)
52- else:52+ def lightning_indexer_metadata_meta(
53- sparse_indices_out = torch.empty([q.shape[0], key_head_num, topk], dtype=torch.int32, device="meta")53+ num_heads_q: int, num_heads_k: int, head_dim: int, topk: int, cu_seqlens_q: Optional[torch.Tensor] = None,
54- if return_value:54+ cu_seqlens_k: Optional[torch.Tensor] = None, seqused_q: Optional[torch.Tensor] = None,
55- if layout_q == "BSND":55+ seqused_k: Optional[torch.Tensor] = None, cmp_residual_k: Optional[torch.Tensor] = None,
56- sparse_values_out = torch.empty([q.shape[0], q.shape[1], key_head_num, topk],56+ batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_k: Optional[int] = None,
57- dtype=torch.float, device="meta")57+ layout_q: Optional[str] = None, layout_k: Optional[str] = None, mask_mode: Optional[int] = None,
58- else:58+ cmp_ratio: Optional[int] = None):
59- sparse_values_out = torch.empty([q.shape[0], key_head_num, topk],59+ return torch.empty((LI_V2_METADATA_SIZE), dtype=torch.int32, device="npu")
60- dtype=torch.float, device="meta")60+ 
61- else:61+ 
62- sparse_values_out = torch.empty([0], dtype=torch.float, device="meta")62+ @impl(AS_LIBRARY, self.name, "Meta")
63- return (sparse_indices_out, sparse_values_out)63+ def lightning_indexer_v2_meta(q, k, w, topk, *, cu_seqlens_q=None, cu_seqlens_k=None,
64- 64+ seqused_q=None, seqused_k=None, cmp_residual_k=None, block_table=None,
65-# Instantiate the builder65+ output_idx_offset=None, metadata=None, max_seqlen_q=-1,
66-npu_lightning_indexer_v2_op_builder = LightningIndexerV2OpBuilder()66+ layout_q="BSND", layout_k="BSND", mask_mode=0,
67-op_module = npu_lightning_indexer_v2_op_builder.load() # Compiles/loads the .so file67+ cmp_ratio=1, return_value=0):
68- 68+ key_head_num = k.shape[1] if layout_k == "TND" else k.shape[2]
69- 69+ 
70-@impl(AS_LIBRARY, npu_lightning_indexer_v2_op_builder.name, "PrivateUse1")70+ if layout_q == "BSND":
71-def npu_lightning_indexer_v2(q, k, w, topk, *, cu_seqlens_q=None, cu_seqlens_k=None,71+ sparse_indices_out = torch.empty([q.shape[0], q.shape[1], key_head_num, topk],
72- seqused_q=None, seqused_k=None, cmp_residual_k=None, block_table=None, 72+ dtype=torch.int32, device="meta")
73- output_idx_offset=None, metadata=None, max_seqlen_q=-1,73+ else:
74- layout_q="BSND", layout_k="BSND", mask_mode=0, 74+ sparse_indices_out = torch.empty([q.shape[0], key_head_num, topk], dtype=torch.int32, device="meta")
75- cmp_ratio=1, return_value=0):75+ if return_value:
76- """76+ if layout_q == "BSND":
77- dispatcher implementation for NPU.zhe77+ sparse_values_out = torch.empty([q.shape[0], q.shape[1], key_head_num, topk],
78- 'PrivateUse1' is the combine key for custom NPU backends.78+ dtype=torch.float, device="meta")
79- """79+ else:
80- return op_module.npu_lightning_indexer_v2(q, k, w, topk, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k, 80+ sparse_values_out = torch.empty([q.shape[0], key_head_num, topk],
81- cmp_residual_k, block_table, output_idx_offset, metadata, max_seqlen_q, 81+ dtype=torch.float, device="meta")
82+ else:
83+ sparse_values_out = torch.empty([0], dtype=torch.float, device="meta")
84+ return (sparse_indices_out, sparse_values_out)
85+ 
86+# Instantiate the builder
87+lightning_indexer_v2_op_builder = LightningIndexerV2OpBuilder()
88+op_module = lightning_indexer_v2_op_builder.load() # Compiles/loads the .so file
89+ 
90+ 
91+@impl(AS_LIBRARY, LI_V2_METADATA_OP_NAME, "PrivateUse1")
92+def lightning_indexer_metadata(
93+ num_heads_q: int, num_heads_k: int, head_dim: int, topk: int, cu_seqlens_q: Optional[torch.Tensor] = None,
94+ cu_seqlens_k: Optional[torch.Tensor] = None, seqused_q: Optional[torch.Tensor] = None,
95+ seqused_k: Optional[torch.Tensor] = None, cmp_residual_k: Optional[torch.Tensor] = None,
96+ batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_k: Optional[int] = None,
97+ layout_q: Optional[str] = None, layout_k: Optional[str] = None, mask_mode: Optional[int] = None,
98+ cmp_ratio: Optional[int] = None):
99+ """
100+ dispatcher implementation for NPU.zhe
101+ 'PrivateUse1' is the combine key for custom NPU backends.
102+ """
103+ batch_size = 0 if batch_size is None else batch_size
104+ max_seqlen_q = -1 if max_seqlen_q is None else max_seqlen_q
105+ max_seqlen_k = -1 if max_seqlen_k is None else max_seqlen_k
106+ layout_q = "BSND" if layout_q is None else layout_q
107+ layout_k = "BSND" if layout_k is None else layout_k
108+ mask_mode = 0 if mask_mode is None else mask_mode
109+ cmp_ratio = 1 if cmp_ratio is None else cmp_ratio
110+ 
111+ return op_module.lightning_indexer_metadata(
112+ num_heads_q, num_heads_k, head_dim, topk, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k, cmp_residual_k,
113+ batch_size, max_seqlen_q, max_seqlen_k, layout_q, layout_k, mask_mode, cmp_ratio)
114+ 
115+ 
116+@torch.library.register_kernel("cann_ops_transformer::" + LI_V2_METADATA_OP_NAME, None)
117+def lightning_indexer_metadata_fallback(
118+ num_heads_q: int, num_heads_k: int, head_dim: int, topk: int, cu_seqlens_q: Optional[torch.Tensor] = None,
119+ cu_seqlens_k: Optional[torch.Tensor] = None, seqused_q: Optional[torch.Tensor] = None,
120+ seqused_k: Optional[torch.Tensor] = None, cmp_residual_k: Optional[torch.Tensor] = None,
121+ batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_k: Optional[int] = None,
122+ layout_q: Optional[str] = None, layout_k: Optional[str] = None, mask_mode: Optional[int] = None,
123+ cmp_ratio: Optional[int] = None):
124+ # 处理所有 tensor 都为 None 的情况
125+ # 调用 NPU 实现
126+ return lightning_indexer_metadata(
127+ num_heads_q, num_heads_k, head_dim, topk, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k, cmp_residual_k,
128+ batch_size, max_seqlen_q, max_seqlen_k, layout_q, layout_k, mask_mode, cmp_ratio)
129+ 
130+torch.compiler.allow_in_graph(lightning_indexer_metadata)
131+ 
132+ 
133+@impl(AS_LIBRARY, lightning_indexer_v2_op_builder.name, "PrivateUse1")
134+def lightning_indexer_v2(q, k, w, topk, *, cu_seqlens_q=None, cu_seqlens_k=None,
135+ seqused_q=None, seqused_k=None, cmp_residual_k=None, block_table=None,
136+ output_idx_offset=None, metadata=None, max_seqlen_q=-1,
137+ layout_q="BSND", layout_k="BSND", mask_mode=0,
138+ cmp_ratio=1, return_value=0):
139+ """
140+ dispatcher implementation for NPU.zhe
141+ 'PrivateUse1' is the combine key for custom NPU backends.
142+ """
143+ return op_module.lightning_indexer_v2(q, k, w, topk, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k,
144+ cmp_residual_k, block_table, output_idx_offset, metadata, max_seqlen_q,
82 layout_q, layout_k, mask_mode, cmp_ratio, return_value)145 layout_q, layout_k, mask_mode, cmp_ratio, return_value)
Rtorch_extension/npu_ops_transformer/ops/mega_moe.pytorch_extension/cann_ops_transformer/ops/mega_moe.py+3-3
@@ -13,8 +13,8 @@ import torch
13import torch_npu13import torch_npu
14from torch.library import impl14from torch.library import impl
15from torch_npu.utils._error_code import ErrCode, ops_error15from torch_npu.utils._error_code import ErrCode, ops_error
16-from npu_ops_transformer.op_builder.builder import OpBuilder16+from cann_ops_transformer.op_builder.builder import OpBuilder
17-from npu_ops_transformer.op_builder.builder import AS_LIBRARY17+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
18from .comm_context import CommContextManager18from .comm_context import CommContextManager
19 19 
20 20 
@@ -204,7 +204,7 @@ def mega_moe(
204 weight2_type: Optional[int] = None,204 weight2_type: Optional[int] = None,
205):205):
206 206 
207- return torch.ops.npu_ops_transformer.npu_mega_moe(207+ return torch.ops.cann_ops_transformer.npu_mega_moe(
208 sym_buffer.context,208 sym_buffer.context,
209 x,209 x,
210 topk_ids,210 topk_ids,
Rtorch_extension/npu_ops_transformer/ops/mhc_post.pytorch_extension/cann_ops_transformer/ops/mhc_post.py+3-3
@@ -10,8 +10,8 @@
10import torch10import torch
11import torch_npu11import torch_npu
12from torch.library import impl12from torch.library import impl
13-from npu_ops_transformer.op_builder.builder import OpBuilder13+from cann_ops_transformer.op_builder.builder import OpBuilder
14-from npu_ops_transformer.op_builder.builder import AS_LIBRARY14+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
15 15 
16 16 
17class MhcPostFunction(torch.autograd.Function):17class MhcPostFunction(torch.autograd.Function):
@@ -24,7 +24,7 @@ class MhcPostFunction(torch.autograd.Function):
24 @staticmethod24 @staticmethod
25 def backward(ctx, grad_output):25 def backward(ctx, grad_output):
26 x, h_res, h_out, h_post = ctx.saved_tensors26 x, h_res, h_out, h_post = ctx.saved_tensors
27- from npu_ops_transformer.ops.mhc_post_backward import mhc_post_backward27+ from cann_ops_transformer.ops.mhc_post_backward import mhc_post_backward
28 grad_x, grad_h_res, grad_h_out, grad_h_post = mhc_post_backward(28 grad_x, grad_h_res, grad_h_out, grad_h_post = mhc_post_backward(
29 grad_output, x, h_res, h_out, h_post29 grad_output, x, h_res, h_out, h_post
30 )30 )
Rtorch_extension/npu_ops_transformer/ops/mhc_post_backward.pytorch_extension/cann_ops_transformer/ops/mhc_post_backward.py+2-2
@@ -10,8 +10,8 @@
10import torch10import torch
11import torch_npu11import torch_npu
12from torch.library import impl12from torch.library import impl
13-from npu_ops_transformer.op_builder.builder import OpBuilder13+from cann_ops_transformer.op_builder.builder import OpBuilder
14-from npu_ops_transformer.op_builder.builder import AS_LIBRARY14+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
15 15 
16 16 
17class MhcPostBackwardOpBuilder(OpBuilder):17class MhcPostBackwardOpBuilder(OpBuilder):
Rtorch_extension/npu_ops_transformer/ops/mhc_pre_sinkhorn.pytorch_extension/cann_ops_transformer/ops/mhc_pre_sinkhorn.py+3-3
@@ -10,8 +10,8 @@
10import torch10import torch
11import torch_npu11import torch_npu
12from torch.library import impl12from torch.library import impl
13-from npu_ops_transformer.op_builder.builder import OpBuilder13+from cann_ops_transformer.op_builder.builder import OpBuilder
14-from npu_ops_transformer.op_builder.builder import AS_LIBRARY14+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
15 15 
16 16 
17class MhcPreSinkhornFunction(torch.autograd.Function):17class MhcPreSinkhornFunction(torch.autograd.Function):
@@ -33,7 +33,7 @@ class MhcPreSinkhornFunction(torch.autograd.Function):
33 x, phi, alpha, bias, h_pre, hc_before_norm, inv_rms, sum_out, norm_out = ctx.saved_tensors33 x, phi, alpha, bias, h_pre, hc_before_norm, inv_rms, sum_out, norm_out = ctx.saved_tensors
34 hc_eps = ctx.hc_eps34 hc_eps = ctx.hc_eps
35 35 
36- from npu_ops_transformer.ops.mhc_pre_sinkhorn_backward import mhc_pre_sinkhorn_backward36+ from cann_ops_transformer.ops.mhc_pre_sinkhorn_backward import mhc_pre_sinkhorn_backward
37 grad_x, grad_phi, grad_alpha, grad_bias = mhc_pre_sinkhorn_backward(37 grad_x, grad_phi, grad_alpha, grad_bias = mhc_pre_sinkhorn_backward(
38 grad_hin, grad_h_post, grad_h_res,38 grad_hin, grad_h_post, grad_h_res,
39 x, phi, alpha, bias,39 x, phi, alpha, bias,
Rtorch_extension/npu_ops_transformer/ops/mhc_pre_sinkhorn_backward.pytorch_extension/cann_ops_transformer/ops/mhc_pre_sinkhorn_backward.py+2-2
@@ -10,8 +10,8 @@
10import torch10import torch
11import torch_npu11import torch_npu
12from torch.library import impl12from torch.library import impl
13-from npu_ops_transformer.op_builder.builder import OpBuilder13+from cann_ops_transformer.op_builder.builder import OpBuilder
14-from npu_ops_transformer.op_builder.builder import AS_LIBRARY14+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
15 15 
16 16 
17class MhcPreSinkhornBackwardOpBuilder(OpBuilder):17class MhcPreSinkhornBackwardOpBuilder(OpBuilder):
Rtorch_extension/npu_ops_transformer/ops/mixed_quant_sparse_flash_mla.pytorch_extension/cann_ops_transformer/ops/mixed_quant_sparse_flash_mla.py+124-26
@@ -6,16 +6,19 @@
6# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.6# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
7# See LICENSE in the root of the software repository for the full text of the License.7# See LICENSE in the root of the software repository for the full text of the License.
8 8 
9+from typing import Optional
9import torch10import torch
10import torch_npu11import torch_npu
11from torch.library import impl12from torch.library import impl
12-from npu_ops_transformer.op_builder.builder import OpBuilder13+from cann_ops_transformer.op_builder.builder import OpBuilder
13-from npu_ops_transformer.op_builder.builder import AS_LIBRARY14+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
15+MQSMLA_METADATA_SIZE = 1024
16+MQSMLA_METADATA_OP_NAME = "mixed_quant_sparse_flash_mla_metadata"
14 17 
15 18 
16class MixedQuantSparseFlashMlaOpBuilder(OpBuilder):19class MixedQuantSparseFlashMlaOpBuilder(OpBuilder):
17 def __init__(self):20 def __init__(self):
18- super(MixedQuantSparseFlashMlaOpBuilder, self).__init__("npu_mixed_quant_sparse_flash_mla")21+ super(MixedQuantSparseFlashMlaOpBuilder, self).__init__("mixed_quant_sparse_flash_mla")
19 22 
20 def sources(self):23 def sources(self):
21 """Path to C++ source code."""24 """Path to C++ source code."""
@@ -23,31 +26,61 @@ class MixedQuantSparseFlashMlaOpBuilder(OpBuilder):
23 26 
24 def schema(self) -> str:27 def schema(self) -> str:
25 """PyTorch operator signature."""28 """PyTorch operator signature."""
26- return "npu_mixed_quant_sparse_flash_mla(Tensor q, *," \29+ return [
27- "Tensor? ori_kv=None, Tensor? cmp_kv=None, " \30+ "mixed_quant_sparse_flash_mla_metadata(int num_heads_q, int num_heads_kv, int head_dim,"
28- "Tensor? ori_sparse_indices=None, Tensor? cmp_sparse_indices=None, " \31+ "int quant_mode, *, Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_ori_kv=None,"
29- "Tensor? ori_block_table=None, Tensor? cmp_block_table=None, " \32+ "Tensor? cu_seqlens_cmp_kv=None, Tensor? seqused_q=None, Tensor? seqused_ori_kv=None,"
30- "Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_ori_kv=None, " \33+ "Tensor? seqused_cmp_kv=None, Tensor? cmp_residual_kv=None, Tensor? ori_topk_length=None,"
31- "Tensor? cu_seqlens_cmp_kv=None, Tensor? seqused_q=None, " \34+ "Tensor? cmp_topk_length=None, int? batch_size=None, int? max_seqlen_q=None, int? max_seqlen_ori_kv=None,"
32- "Tensor? seqused_ori_kv=None, Tensor? seqused_cmp_kv=None, " \35+ "int? max_seqlen_cmp_kv=None, int? ori_topk=None, int? cmp_topk=None, int? rope_head_dim=None,"
33- "Tensor? cmp_residual_kv=None, " \36+ "int? cmp_ratio=None, int? ori_mask_mode=None, int? cmp_mask_mode=None, int? ori_win_left=None,"
34- "Tensor? ori_topk_length=None, Tensor? cmp_topk_length=None, " \37+ "int? ori_win_right=None, str? layout_q=None, str? layout_kv=None, bool? has_ori_kv=None,"
35- "Tensor? sinks=None, Tensor? metadata=None, " \38+ "bool? has_cmp_kv=None) -> Tensor",
36- "int quant_mode=None, int rope_head_dim=None, " \39+ 
37- "float softmax_scale=None, int cmp_ratio=None, " \40+ "mixed_quant_sparse_flash_mla(Tensor q, *,"
38- "int ori_mask_mode=0, int cmp_mask_mode=0, " \41+ "Tensor? ori_kv=None, Tensor? cmp_kv=None, "
39- "int ori_win_left=-1, int ori_win_right=-1, " \42+ "Tensor? ori_sparse_indices=None, Tensor? cmp_sparse_indices=None, "
40- "str layout_q=\"BSND\", str layout_kv=\"BSND\", " \43+ "Tensor? ori_block_table=None, Tensor? cmp_block_table=None, "
41- "int topk_value_mode=1, bool return_softmax_lse=False, " \44+ "Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_ori_kv=None, "
42- "int? key_dtype=None, int? value_dtype=None) -> (Tensor, Tensor)"45+ "Tensor? cu_seqlens_cmp_kv=None, Tensor? seqused_q=None, "
46+ "Tensor? seqused_ori_kv=None, Tensor? seqused_cmp_kv=None, "
47+ "Tensor? cmp_residual_kv=None, "
48+ "Tensor? ori_topk_length=None, Tensor? cmp_topk_length=None, "
49+ "Tensor? sinks=None, Tensor? metadata=None, "
50+ "int quant_mode=None, int rope_head_dim=None, "
51+ "float softmax_scale=None, int cmp_ratio=None, "
52+ "int ori_mask_mode=0, int cmp_mask_mode=0, "
53+ "int ori_win_left=-1, int ori_win_right=-1, "
54+ "str layout_q=\"BSND\", str layout_kv=\"BSND\", "
55+ "int topk_value_mode=1, bool return_softmax_lse=False, "
56+ "int? key_dtype=None, int? value_dtype=None) -> (Tensor, Tensor)"
57+ ]
58+ 
43 59 
44 def register_meta(self):60 def register_meta(self):
45 """61 """
46 Registers the Meta implementation (Shape/Dtype inference).62 Registers the Meta implementation (Shape/Dtype inference).
47 Essential for Autograd and FakeTensor support.63 Essential for Autograd and FakeTensor support.
48 """64 """
65+ @torch.library.register_fake("cann_ops_transformer::" + MQSMLA_METADATA_OP_NAME)
66+ def mixed_quant_sparse_flash_mla_metadata_meta(
67+ num_heads_q: int, num_heads_kv: int, head_dim: int, quant_mode: int,
68+ cu_seqlens_q: Optional[torch.Tensor] = None, cu_seqlens_ori_kv: Optional[torch.Tensor] = None,
69+ cu_seqlens_cmp_kv: Optional[torch.Tensor] = None, seqused_q: Optional[torch.Tensor] = None,
70+ seqused_ori_kv: Optional[torch.Tensor] = None, seqused_cmp_kv: Optional[torch.Tensor] = None,
71+ cmp_residual_kv: Optional[torch.Tensor] = None, ori_topk_length: Optional[torch.Tensor] = None,
72+ cmp_topk_length: Optional[torch.Tensor] = None, batch_size: Optional[int] = None,
73+ max_seqlen_q: Optional[int] = None, max_seqlen_ori_kv: Optional[int] = None,
74+ max_seqlen_cmp_kv: Optional[int] = None, ori_topk: Optional[int] = None, cmp_topk: Optional[int] = None,
75+ rope_head_dim: Optional[int] = None, cmp_ratio: Optional[int] = None, ori_mask_mode: Optional[int] = None,
76+ cmp_mask_mode: Optional[int] = None, ori_win_left: Optional[int] = None,
77+ ori_win_right: Optional[int] = None, layout_q: Optional[str] = None, layout_kv: Optional[str] = None,
78+ has_ori_kv: Optional[bool] = None, has_cmp_kv: Optional[bool] = None):
79+ return torch.empty((MQSMLA_METADATA_SIZE), dtype=torch.int32, device="npu")
80+ 
81+ 
49 @impl(AS_LIBRARY, self.name, "Meta")82 @impl(AS_LIBRARY, self.name, "Meta")
50- def npu_mixed_quant_sparse_flash_mla_meta(q,83+ def mixed_quant_sparse_flash_mla_meta(q,
51 ori_kv=None, cmp_kv=None,84 ori_kv=None, cmp_kv=None,
52 ori_sparse_indices=None, cmp_sparse_indices=None,85 ori_sparse_indices=None, cmp_sparse_indices=None,
53 ori_block_table=None, cmp_block_table=None,86 ori_block_table=None, cmp_block_table=None,
@@ -84,12 +117,77 @@ class MixedQuantSparseFlashMlaOpBuilder(OpBuilder):
84 return (attn_out, softmax_lse)117 return (attn_out, softmax_lse)
85 118 
86# Instantiate the builder119# Instantiate the builder
87-npu_mixed_quant_sparse_flash_mla_op_builder = MixedQuantSparseFlashMlaOpBuilder()120+mixed_quant_sparse_flash_mla_op_builder = MixedQuantSparseFlashMlaOpBuilder()
88-op_module = npu_mixed_quant_sparse_flash_mla_op_builder.load() # Compiles/loads the .so file121+op_module = mixed_quant_sparse_flash_mla_op_builder.load() # Compiles/loads the .so file
89 122 
90 123 
91-@impl(AS_LIBRARY, npu_mixed_quant_sparse_flash_mla_op_builder.name, "PrivateUse1")124+@impl(AS_LIBRARY, MQSMLA_METADATA_OP_NAME, "PrivateUse1")
92-def npu_mixed_quant_sparse_flash_mla(q,125+def mixed_quant_sparse_flash_mla_metadata(
126+ num_heads_q: int, num_heads_kv: int, head_dim: int, quant_mode: int, cu_seqlens_q: Optional[torch.Tensor] = None,
127+ cu_seqlens_ori_kv: Optional[torch.Tensor] = None, cu_seqlens_cmp_kv: Optional[torch.Tensor] = None,
128+ seqused_q: Optional[torch.Tensor] = None, seqused_ori_kv: Optional[torch.Tensor] = None,
129+ seqused_cmp_kv: Optional[torch.Tensor] = None, cmp_residual_kv: Optional[torch.Tensor] = None,
130+ ori_topk_length: Optional[torch.Tensor] = None, cmp_topk_length: Optional[torch.Tensor] = None,
131+ batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_ori_kv: Optional[int] = None,
132+ max_seqlen_cmp_kv: Optional[int] = None, ori_topk: Optional[int] = None, cmp_topk: Optional[int] = None,
133+ rope_head_dim: Optional[int] = None, cmp_ratio: Optional[int] = None, ori_mask_mode: Optional[int] = None,
134+ cmp_mask_mode: Optional[int] = None, ori_win_left: Optional[int] = None, ori_win_right: Optional[int] = None,
135+ layout_q: Optional[str] = None, layout_kv: Optional[str] = None, has_ori_kv: Optional[bool] = None,
136+ has_cmp_kv: Optional[bool] = None):
137+ """
138+ Dispatcher implementation: NPU.
139+ 'PrivateUse1' is dispatch key for custom NPU backends.
140+ """
141+ batch_size = 0 if batch_size is None else batch_size
142+ max_seqlen_q = 0 if max_seqlen_q is None else max_seqlen_q
143+ max_seqlen_ori_kv = 0 if max_seqlen_ori_kv is None else max_seqlen_ori_kv
144+ max_seqlen_cmp_kv = 0 if max_seqlen_cmp_kv is None else max_seqlen_cmp_kv
145+ ori_topk = 0 if ori_topk is None else ori_topk
146+ cmp_topk = 0 if cmp_topk is None else cmp_topk
147+ rope_head_dim = 64 if rope_head_dim is None else rope_head_dim
148+ cmp_ratio = 1 if cmp_ratio is None else cmp_ratio
149+ ori_mask_mode = 0 if ori_mask_mode is None else ori_mask_mode
150+ cmp_mask_mode = 0 if cmp_mask_mode is None else cmp_mask_mode
151+ ori_win_left = -1 if ori_win_left is None else ori_win_left
152+ ori_win_right = -1 if ori_win_right is None else ori_win_right
153+ layout_q = "BSND" if layout_q is None else layout_q
154+ layout_kv = "BSND" if layout_kv is None else layout_kv
155+ has_ori_kv = True if has_ori_kv is None else has_ori_kv
156+ has_cmp_kv = True if has_cmp_kv is None else has_cmp_kv
157+ 
158+ return op_module.mixed_quant_sparse_flash_mla_metadata(
159+ num_heads_q, num_heads_kv, head_dim, quant_mode, cu_seqlens_q, cu_seqlens_ori_kv, cu_seqlens_cmp_kv, seqused_q,
160+ seqused_ori_kv, seqused_cmp_kv, cmp_residual_kv, ori_topk_length, cmp_topk_length, batch_size, max_seqlen_q,
161+ max_seqlen_ori_kv, max_seqlen_cmp_kv, ori_topk, cmp_topk, rope_head_dim, cmp_ratio, ori_mask_mode,
162+ cmp_mask_mode, ori_win_left, ori_win_right, layout_q, layout_kv, has_ori_kv, has_cmp_kv)
163+ 
164+ 
165+@torch.library.register_kernel("cann_ops_transformer::" + MQSMLA_METADATA_OP_NAME, None)
166+def mixed_quant_sparse_flash_mla_metadata_fallback(
167+ num_heads_q: int, num_heads_kv: int, head_dim: int, quant_mode: int, cu_seqlens_q: Optional[torch.Tensor] = None,
168+ cu_seqlens_ori_kv: Optional[torch.Tensor] = None, cu_seqlens_cmp_kv: Optional[torch.Tensor] = None,
169+ seqused_q: Optional[torch.Tensor] = None, seqused_ori_kv: Optional[torch.Tensor] = None,
170+ seqused_cmp_kv: Optional[torch.Tensor] = None, cmp_residual_kv: Optional[torch.Tensor] = None,
171+ ori_topk_length: Optional[torch.Tensor] = None, cmp_topk_length: Optional[torch.Tensor] = None,
172+ batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_ori_kv: Optional[int] = None,
173+ max_seqlen_cmp_kv: Optional[int] = None, ori_topk: Optional[int] = None, cmp_topk: Optional[int] = None,
174+ rope_head_dim: Optional[int] = None, cmp_ratio: Optional[int] = None, ori_mask_mode: Optional[int] = None,
175+ cmp_mask_mode: Optional[int] = None, ori_win_left: Optional[int] = None, ori_win_right: Optional[int] = None,
176+ layout_q: Optional[str] = None, layout_kv: Optional[str] = None, has_ori_kv: Optional[bool] = None,
177+ has_cmp_kv: Optional[bool] = None):
178+ # 处理所有 tensor 都为 None 的情况
179+ # 调用 NPU 实现
180+ return mixed_quant_sparse_flash_mla_metadata(
181+ num_heads_q, num_heads_kv, head_dim, quant_mode, cu_seqlens_q, cu_seqlens_ori_kv, cu_seqlens_cmp_kv, seqused_q,
182+ seqused_ori_kv, seqused_cmp_kv, cmp_residual_kv, ori_topk_length, cmp_topk_length, batch_size, max_seqlen_q,
183+ max_seqlen_ori_kv, max_seqlen_cmp_kv, ori_topk, cmp_topk, rope_head_dim, cmp_ratio, ori_mask_mode,
184+ cmp_mask_mode, ori_win_left, ori_win_right, layout_q, layout_kv, has_ori_kv, has_cmp_kv)
185+ 
186+torch.compiler.allow_in_graph(mixed_quant_sparse_flash_mla_metadata)
187+ 
188+ 
189+@impl(AS_LIBRARY, mixed_quant_sparse_flash_mla_op_builder.name, "PrivateUse1")
190+def mixed_quant_sparse_flash_mla(q,
93 ori_kv=None, cmp_kv=None,191 ori_kv=None, cmp_kv=None,
94 ori_sparse_indices=None, cmp_sparse_indices=None,192 ori_sparse_indices=None, cmp_sparse_indices=None,
95 ori_block_table=None, cmp_block_table=None,193 ori_block_table=None, cmp_block_table=None,
@@ -110,7 +208,7 @@ def npu_mixed_quant_sparse_flash_mla(q,
110 dispatcher implementation for NPU208 dispatcher implementation for NPU
111 'PrivateUse1' is the combine key for custom NPU backends.209 'PrivateUse1' is the combine key for custom NPU backends.
112 """210 """
113- return op_module.npu_mixed_quant_sparse_flash_mla(q,211+ return op_module.mixed_quant_sparse_flash_mla(q,
114 ori_kv, cmp_kv,212 ori_kv, cmp_kv,
115 ori_sparse_indices, cmp_sparse_indices,213 ori_sparse_indices, cmp_sparse_indices,
116 ori_block_table, cmp_block_table,214 ori_block_table, cmp_block_table,
Rtorch_extension/npu_ops_transformer/ops/moe_distribute_combine.pytorch_extension/cann_ops_transformer/ops/moe_distribute_combine.py+2-2
@@ -10,8 +10,8 @@
10import torch10import torch
11import torch_npu11import torch_npu
12from torch.library import impl12from torch.library import impl
13-from npu_ops_transformer.op_builder.builder import OpBuilder13+from cann_ops_transformer.op_builder.builder import OpBuilder
14-from npu_ops_transformer.op_builder.builder import AS_LIBRARY14+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
15 15 
16 16 
17class MoeDistributeCombineOpBuilder(OpBuilder):17class MoeDistributeCombineOpBuilder(OpBuilder):
Rtorch_extension/npu_ops_transformer/ops/moe_distribute_dispatch.pytorch_extension/cann_ops_transformer/ops/moe_distribute_dispatch.py+2-2
@@ -12,8 +12,8 @@ import torch
12import torch_npu12import torch_npu
13from torch.library import impl13from torch.library import impl
14from torch_npu.utils._error_code import ErrCode, ops_error14from torch_npu.utils._error_code import ErrCode, ops_error
15-from npu_ops_transformer.op_builder.builder import OpBuilder15+from cann_ops_transformer.op_builder.builder import OpBuilder
16-from npu_ops_transformer.op_builder.builder import AS_LIBRARY16+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
17 17 
18 18 
19TORCH_DTYPE_ENUM_VALUE_TO_SCALAR_TYPE_MAP = {19TORCH_DTYPE_ENUM_VALUE_TO_SCALAR_TYPE_MAP = {
Rtorch_extension/npu_ops_transformer/ops/npu_sparse_lightning_indexer_kl_loss_grad.pytorch_extension/cann_ops_transformer/ops/npu_sparse_lightning_indexer_kl_loss_grad.py+3-3
@@ -12,8 +12,8 @@ import torch
12import torch_npu12import torch_npu
13from torch.library import impl13from torch.library import impl
14 14 
15-from npu_ops_transformer.op_builder.builder import AS_LIBRARY15+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
16-from npu_ops_transformer.op_builder.builder import OpBuilder16+from cann_ops_transformer.op_builder.builder import OpBuilder
17 17 
18 18 
19class SparseLightningIndexerKLLossGradOpBuilder(OpBuilder):19class SparseLightningIndexerKLLossGradOpBuilder(OpBuilder):
@@ -117,7 +117,7 @@ except ImportError:
117 117 
118if _TORCHAIR_AVAILABLE:118if _TORCHAIR_AVAILABLE:
119 @register_fx_node_ge_converter(119 @register_fx_node_ge_converter(
120- torch.ops.npu_ops_transformer.npu_sparse_lightning_indexer_kl_loss_grad.default120+ torch.ops.cann_ops_transformer.npu_sparse_lightning_indexer_kl_loss_grad.default
121 )121 )
122 def convert_npu_sparse_lightning_indexer_kl_loss_grad(122 def convert_npu_sparse_lightning_indexer_kl_loss_grad(
123 q: Tensor,123 q: Tensor,
Rtorch_extension/npu_ops_transformer/ops/npu_sparse_lightning_indexer_kl_loss_grad_metadata.pytorch_extension/cann_ops_transformer/ops/npu_sparse_lightning_indexer_kl_loss_grad_metadata.py+4-4
@@ -12,8 +12,8 @@ import torch
12import torch_npu12import torch_npu
13from torch.library import impl13from torch.library import impl
14 14 
15-from npu_ops_transformer.op_builder.builder import AS_LIBRARY15+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
16-from npu_ops_transformer.op_builder.builder import OpBuilder16+from cann_ops_transformer.op_builder.builder import OpBuilder
17 17 
18 18 
19SLI_KL_LOSS_GRAD_METADATA_SIZE = 6419SLI_KL_LOSS_GRAD_METADATA_SIZE = 64
@@ -118,7 +118,7 @@ def npu_sparse_lightning_indexer_kl_loss_grad_metadata(
118 118 
119 119 
120@torch.library.register_kernel(120@torch.library.register_kernel(
121- "npu_ops_transformer::npu_sparse_lightning_indexer_kl_loss_grad_metadata", None121+ "cann_ops_transformer::npu_sparse_lightning_indexer_kl_loss_grad_metadata", None
122)122)
123def npu_sparse_lightning_indexer_kl_loss_grad_metadata_fallback(123def npu_sparse_lightning_indexer_kl_loss_grad_metadata_fallback(
124 num_heads_q,124 num_heads_q,
@@ -172,7 +172,7 @@ except ImportError:
172 172 
173if _TORCHAIR_AVAILABLE:173if _TORCHAIR_AVAILABLE:
174 @register_fx_node_ge_converter(174 @register_fx_node_ge_converter(
175- torch.ops.npu_ops_transformer.npu_sparse_lightning_indexer_kl_loss_grad_metadata.default175+ torch.ops.cann_ops_transformer.npu_sparse_lightning_indexer_kl_loss_grad_metadata.default
176 )176 )
177 def convert_npu_sparse_lightning_indexer_kl_loss_grad_metadata(177 def convert_npu_sparse_lightning_indexer_kl_loss_grad_metadata(
178 num_heads_q: int,178 num_heads_q: int,
Rtorch_extension/npu_ops_transformer/ops/quant_lightning_indexer_v2.pytorch_extension/cann_ops_transformer/ops/quant_lightning_indexer_v2.py+78-14
@@ -7,16 +7,19 @@
7# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.7# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8# See LICENSE in the root of the software repository for the full text of the License.8# See LICENSE in the root of the software repository for the full text of the License.
9# -----------------------------------------------------------------------------------------------------------9# -----------------------------------------------------------------------------------------------------------
10+from typing import Optional
10import torch11import torch
11import torch_npu12import torch_npu
12from torch.library import impl13from torch.library import impl
13-from npu_ops_transformer.op_builder.builder import OpBuilder14+from cann_ops_transformer.op_builder.builder import OpBuilder
14-from npu_ops_transformer.op_builder.builder import AS_LIBRARY15+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
16+QLI_V2_METADATA_SIZE = 1024
17+QLI_V2_METADATA_OP_NAME = "quant_lightning_indexer_metadata"
15 18 
16 19 
17class QuantLightningIndexerV2OpBuilder(OpBuilder):20class QuantLightningIndexerV2OpBuilder(OpBuilder):
18 def __init__(self):21 def __init__(self):
19- super(QuantLightningIndexerV2OpBuilder, self).__init__("npu_quant_lightning_indexer_v2")22+ super(QuantLightningIndexerV2OpBuilder, self).__init__("quant_lightning_indexer_v2")
20 23 
21 def sources(self):24 def sources(self):
22 """Path to C++ source code."""25 """Path to C++ source code."""
@@ -24,20 +27,39 @@ class QuantLightningIndexerV2OpBuilder(OpBuilder):
24 27 
25 def schema(self) -> str:28 def schema(self) -> str:
26 """PyTorch operator signature."""29 """PyTorch operator signature."""
27- return "npu_quant_lightning_indexer_v2(Tensor query, Tensor key, Tensor weights, Tensor query_dequant_scale, "\30+ return [
28- "Tensor key_dequant_scale, int topk, int quant_mode, *, Tensor? cu_seqlens_q=None, "\31+ "quant_lightning_indexer_metadata(int num_heads_q, int num_heads_k, int head_dim, int topk, "
29- "Tensor? cu_seqlens_k=None, Tensor? seqused_q=None, Tensor? seqused_k=None, Tensor? "\32+ "int quant_mode, *, Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_k=None, Tensor? seqused_q=None,"
30- "cmp_residual_k = None, Tensor? block_table=None, Tensor? output_idx_offset=None, Tensor? metadata=None, "\33+ "Tensor? seqused_k=None, Tensor? cmp_residual_k=None, int? batch_size=None, int? max_seqlen_q=None,"
31- "int max_seqlen_q=-1, str layout_q=\"BSND\", str layout_k=\"BSND\", int mask_mode=0, "\34+ "int? max_seqlen_k=None, str? layout_q=None, str? layout_k=None, int? mask_mode=None, "
35+ "int? cmp_ratio=None) -> Tensor",
36+ 
37+ "quant_lightning_indexer_v2(Tensor query, Tensor key, Tensor weights, Tensor query_dequant_scale, "
38+ "Tensor key_dequant_scale, int topk, int quant_mode, *, Tensor? cu_seqlens_q=None, "
39+ "Tensor? cu_seqlens_k=None, Tensor? seqused_q=None, Tensor? seqused_k=None, Tensor? "
40+ "cmp_residual_k = None, Tensor? block_table=None, Tensor? output_idx_offset=None, Tensor? metadata=None, "
41+ "int max_seqlen_q=-1, str layout_q=\"BSND\", str layout_k=\"BSND\", int mask_mode=0, "
32 "int cmp_ratio=1, int return_value=0) -> (Tensor, Tensor)"42 "int cmp_ratio=1, int return_value=0) -> (Tensor, Tensor)"
43+ ]
33 44 
34 def register_meta(self):45 def register_meta(self):
35 """46 """
36 Registers the Meta implementation (Shape/Dtype inference).47 Registers the Meta implementation (Shape/Dtype inference).
37 Essential for Autograd and FakeTensor support.48 Essential for Autograd and FakeTensor support.
38 """49 """
50+ @torch.library.register_fake("cann_ops_transformer::" + QLI_V2_METADATA_OP_NAME)
51+ def quant_lightning_indexer_metadata_meta(
52+ num_heads_q: int, num_heads_k: int, head_dim: int, topk: int, quant_mode: int,
53+ cu_seqlens_q: Optional[torch.Tensor] = None, cu_seqlens_k: Optional[torch.Tensor] = None,
54+ seqused_q: Optional[torch.Tensor] = None, seqused_k: Optional[torch.Tensor] = None,
55+ cmp_residual_k: Optional[torch.Tensor] = None, batch_size: Optional[int] = None,
56+ max_seqlen_q: Optional[int] = None, max_seqlen_k: Optional[int] = None, layout_q: Optional[str] = None,
57+ layout_k: Optional[str] = None, mask_mode: Optional[int] = None, cmp_ratio: Optional[int] = None):
58+ return torch.empty((QLI_V2_METADATA_SIZE), dtype=torch.int32, device="npu")
59+ 
60+ 
39 @impl(AS_LIBRARY, self.name, "Meta")61 @impl(AS_LIBRARY, self.name, "Meta")
40- def npu_quant_lightning_indexer_v2_meta(query, key, weights, query_dequant_scale, key_dequant_scale, topk,62+ def quant_lightning_indexer_v2_meta(query, key, weights, query_dequant_scale, key_dequant_scale, topk,
41 quant_mode, *, cu_seqlens_q=None, cu_seqlens_k=None,63 quant_mode, *, cu_seqlens_q=None, cu_seqlens_k=None,
42 block_table=None, output_idx_offset=None, metadata=None, 64 block_table=None, output_idx_offset=None, metadata=None,
43 max_seqlen_q=-1, layout_q="BSND", layout_k="BSND",65 max_seqlen_q=-1, layout_q="BSND", layout_k="BSND",
@@ -66,12 +88,54 @@ class QuantLightningIndexerV2OpBuilder(OpBuilder):
66 return (sparse_indices_out, sparse_values_out)88 return (sparse_indices_out, sparse_values_out)
67 89 
68# Instantiate the builder90# Instantiate the builder
69-npu_quant_lightning_indexer_v2_op_builder = QuantLightningIndexerV2OpBuilder()91+quant_lightning_indexer_v2_op_builder = QuantLightningIndexerV2OpBuilder()
70-op_module = npu_quant_lightning_indexer_v2_op_builder.load() # Compiles/loads the .so file92+op_module = quant_lightning_indexer_v2_op_builder.load() # Compiles/loads the .so file
71 93 
72 94 
73-@impl(AS_LIBRARY, npu_quant_lightning_indexer_v2_op_builder.name, "PrivateUse1")95+@impl(AS_LIBRARY, QLI_V2_METADATA_OP_NAME, "PrivateUse1")
74-def npu_quant_lightning_indexer_v2(query, key, weights, query_dequant_scale, key_dequant_scale, topk,96+def quant_lightning_indexer_metadata(
97+ num_heads_q: int, num_heads_k: int, head_dim: int, topk: int, quant_mode: int,
98+ cu_seqlens_q: Optional[torch.Tensor] = None, cu_seqlens_k: Optional[torch.Tensor] = None,
99+ seqused_q: Optional[torch.Tensor] = None, seqused_k: Optional[torch.Tensor] = None,
100+ cmp_residual_k: Optional[torch.Tensor] = None, batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None,
101+ max_seqlen_k: Optional[int] = None, layout_q: Optional[str] = None, layout_k: Optional[str] = None,
102+ mask_mode: Optional[int] = None, cmp_ratio: Optional[int] = None):
103+ """
104+ dispatcher implementation for NPU.zhe
105+ 'PrivateUse1' is the combine key for custom NPU backends.
106+ """
107+ batch_size = 0 if batch_size is None else batch_size
108+ max_seqlen_q = -1 if max_seqlen_q is None else max_seqlen_q
109+ max_seqlen_k = -1 if max_seqlen_k is None else max_seqlen_k
110+ layout_q = "BSND" if layout_q is None else layout_q
111+ layout_k = "BSND" if layout_k is None else layout_k
112+ mask_mode = 0 if mask_mode is None else mask_mode
113+ cmp_ratio = 1 if cmp_ratio is None else cmp_ratio
114+ 
115+ return op_module.quant_lightning_indexer_metadata(
116+ num_heads_q, num_heads_k, head_dim, topk, quant_mode, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k,
117+ cmp_residual_k, batch_size, max_seqlen_q, max_seqlen_k, layout_q, layout_k, mask_mode, cmp_ratio)
118+ 
119+ 
120+@torch.library.register_kernel("cann_ops_transformer::" + QLI_V2_METADATA_OP_NAME, None)
121+def quant_lightning_indexer_metadata_fallback(
122+ num_heads_q: int, num_heads_k: int, head_dim: int, topk: int, quant_mode: int,
123+ cu_seqlens_q: Optional[torch.Tensor] = None, cu_seqlens_k: Optional[torch.Tensor] = None,
124+ seqused_q: Optional[torch.Tensor] = None, seqused_k: Optional[torch.Tensor] = None,
125+ cmp_residual_k: Optional[torch.Tensor] = None, batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None,
126+ max_seqlen_k: Optional[int] = None, layout_q: Optional[str] = None, layout_k: Optional[str] = None,
127+ mask_mode: Optional[int] = None, cmp_ratio: Optional[int] = None):
128+ # 处理所有 tensor 都为 None 的情况
129+ # 调用 NPU 实现
130+ return quant_lightning_indexer_metadata(
131+ num_heads_q, num_heads_k, head_dim, topk, quant_mode, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k,
132+ cmp_residual_k, batch_size, max_seqlen_q, max_seqlen_k, layout_q, layout_k, mask_mode, cmp_ratio)
133+ 
134+torch.compiler.allow_in_graph(quant_lightning_indexer_metadata)
135+ 
136+ 
137+@impl(AS_LIBRARY, quant_lightning_indexer_v2_op_builder.name, "PrivateUse1")
138+def quant_lightning_indexer_v2(query, key, weights, query_dequant_scale, key_dequant_scale, topk,
75 quant_mode, *, cu_seqlens_q=None, cu_seqlens_k=None,139 quant_mode, *, cu_seqlens_q=None, cu_seqlens_k=None,
76 seqused_q=None, seqused_k=None, cmp_residual_k=None, 140 seqused_q=None, seqused_k=None, cmp_residual_k=None,
77 block_table=None, output_idx_offset=None, metadata=None, 141 block_table=None, output_idx_offset=None, metadata=None,
@@ -81,7 +145,7 @@ def npu_quant_lightning_indexer_v2(query, key, weights, query_dequant_scale, key
81 dispatcher implementation for NPU.zhe145 dispatcher implementation for NPU.zhe
82 'PrivateUse1' is the combine key for custom NPU backends.146 'PrivateUse1' is the combine key for custom NPU backends.
83 """147 """
84- return op_module.npu_quant_lightning_indexer_v2(query, key, weights, query_dequant_scale, key_dequant_scale, 148+ return op_module.quant_lightning_indexer_v2(query, key, weights, query_dequant_scale, key_dequant_scale,
85 topk, quant_mode, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k,149 topk, quant_mode, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k,
86 cmp_residual_k, block_table, output_idx_offset, metadata,150 cmp_residual_k, block_table, output_idx_offset, metadata,
87 max_seqlen_q, layout_q, layout_k, mask_mode,151 max_seqlen_q, layout_q, layout_k, mask_mode,
@@ -0,0 +1,217 @@
1+# -----------------------------------------------------------------------------------------------------------
2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+# CANN Open Software License Agreement Version 2.0 (the "License").
5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+# See LICENSE in the root of the software repository for the full text of the License.
9+# -----------------------------------------------------------------------------------------------------------
10+ 
11+from typing import Optional
12+import torch
13+import torch_npu
14+from torch.library import impl
15+from cann_ops_transformer.op_builder.builder import OpBuilder
16+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
17+ 
18+SMLA_METADATA_SIZE = 1024
19+SMLA_METADATA_OP_NAME = "sparse_flash_mla_metadata"
20+ 
21+ 
22+class SparseFlashMlaOpBuilder(OpBuilder):
23+ def __init__(self):
24+ super(SparseFlashMlaOpBuilder, self).__init__("sparse_flash_mla")
25+ 
26+ def sources(self):
27+ """Path to C++ source code."""
28+ return ['ops/csrc/sparse_flash_mla.cpp']
29+ 
30+ def schema(self) -> str:
31+ """PyTorch operator signature."""
32+ return [
33+ "sparse_flash_mla_metadata(int num_heads_q, int num_heads_kv, int head_dim, *, "
34+ "Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_ori_kv=None, Tensor? cu_seqlens_cmp_kv=None, "
35+ "Tensor? seqused_q=None, Tensor? seqused_ori_kv=None, Tensor? seqused_cmp_kv=None, "
36+ "Tensor? cmp_residual_kv=None, Tensor? ori_topk_length=None, Tensor? cmp_topk_length=None, "
37+ "int? batch_size=None, int? max_seqlen_q=None, int? max_seqlen_ori_kv=None, int? max_seqlen_cmp_kv=None,"
38+ "int? ori_topk=None, int? cmp_topk=None, int? cmp_ratio=None, int? ori_mask_mode=None,"
39+ "int? cmp_mask_mode=None, int? ori_win_left=None, int? ori_win_right=None, str? layout_q=None,"
40+ "str? layout_kv=None, bool? has_ori_kv=None, bool? has_cmp_kv=None) -> Tensor",
41+ 
42+ "sparse_flash_mla(Tensor q, *,"
43+ "Tensor? ori_kv=None, Tensor? cmp_kv=None, "
44+ "Tensor? ori_sparse_indices=None, Tensor? cmp_sparse_indices=None, "
45+ "Tensor? ori_block_table=None, Tensor? cmp_block_table=None, "
46+ "Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_ori_kv=None, "
47+ "Tensor? cu_seqlens_cmp_kv=None, Tensor? seqused_q=None, "
48+ "Tensor? seqused_ori_kv=None, Tensor? seqused_cmp_kv=None, "
49+ "Tensor? cmp_residual_kv=None, "
50+ "Tensor? ori_topk_length=None, Tensor? cmp_topk_length=None, "
51+ "Tensor? sinks=None, Tensor? metadata=None, "
52+ "float softmax_scale=1.0, int cmp_ratio=1, "
53+ "int ori_mask_mode=4, int cmp_mask_mode=3, "
54+ "int ori_win_left=127, int ori_win_right=0, "
55+ "str layout_q=\"BSND\", str layout_kv=\"PA_BBND\", "
56+ "int topk_value_mode=1, bool return_softmax_lse=False) -> (Tensor, Tensor)"
57+ ]
58+
59+ def register_meta(self):
60+ """
61+ Registers the Meta implementation (Shape/Dtype inference).
62+ Essential for Autograd and FakeTensor support.
63+ """
64+ @torch.library.register_fake("cann_ops_transformer::" + SMLA_METADATA_OP_NAME)
65+ def sparse_flash_mla_metadata_meta(
66+ num_heads_q: int, num_heads_kv: int, head_dim: int, cu_seqlens_q: Optional[torch.Tensor] = None,
67+ cu_seqlens_ori_kv: Optional[torch.Tensor] = None, cu_seqlens_cmp_kv: Optional[torch.Tensor] = None,
68+ seqused_q: Optional[torch.Tensor] = None, seqused_ori_kv: Optional[torch.Tensor] = None,
69+ seqused_cmp_kv: Optional[torch.Tensor] = None, cmp_residual_kv: Optional[torch.Tensor] = None,
70+ ori_topk_length: Optional[torch.Tensor] = None, cmp_topk_length: Optional[torch.Tensor] = None,
71+ batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None,
72+ max_seqlen_ori_kv: Optional[int] = None, max_seqlen_cmp_kv: Optional[int] = None,
73+ ori_topk: Optional[int] = None, cmp_topk: Optional[int] = None, cmp_ratio: Optional[int] = None,
74+ ori_mask_mode: Optional[int] = None, cmp_mask_mode: Optional[int] = None,
75+ ori_win_left: Optional[int] = None, ori_win_right: Optional[int] = None, layout_q: Optional[str] = None,
76+ layout_kv: Optional[str] = None, has_ori_kv: Optional[bool] = None, has_cmp_kv: Optional[bool] = None):
77+ return torch.empty((SMLA_METADATA_SIZE), dtype=torch.int32, device="npu")
78+ 
79+ 
80+ @impl(AS_LIBRARY, self.name, "Meta")
81+ def sparse_flash_mla_meta(q,
82+ ori_kv=None, cmp_kv=None,
83+ ori_sparse_indices=None, cmp_sparse_indices=None,
84+ ori_block_table=None, cmp_block_table=None,
85+ cu_seqlens_q=None, cu_seqlens_ori_kv=None,
86+ cu_seqlens_cmp_kv=None, seqused_q=None,
87+ seqused_ori_kv=None, seqused_cmp_kv=None,
88+ cmp_residual_kv=None,
89+ ori_topk_length=None, cmp_topk_length=None,
90+ sinks=None, metadata=None,
91+ softmax_scale=1.0, cmp_ratio=1,
92+ ori_mask_mode=4, cmp_mask_mode=3,
93+ ori_win_left=127, ori_win_right=0,
94+ layout_q='BSND', layout_kv='PA_BBND',
95+ topk_value_mode=1, return_softmax_lse=False):
96+ key_headnum = ori_kv.shape[1] if layout_kv == "TND" else ori_kv.shape[2]
97+ if layout_q == "BSND":
98+ ## 添加softmax_lse
99+ attn_out = torch.empty(q.shape, dtype=q.dtype, device="meta")
100+ if return_softmax_lse:
101+ softmax_lse = torch.empty([q.shape[0], ori_kv.shape[2], q.shape[1], q.shape[2] / ori_kv.shape[2]],
102+ dtype=torch.float32, device="meta")
103+ else:
104+ # 给一个空的合法张量,不能是 nullptr
105+ softmax_lse = torch.empty([], dtype=torch.float32, device="meta")
106+ else:
107+ attn_out = torch.empty(q.shape, dtype=q.dtype, device="meta")
108+ if return_softmax_lse:
109+ softmax_lse = torch.empty([ori_kv.shape[1], q.shape[0], q.shape[1] / ori_kv.shape[1]],
110+ dtype=torch.float32, device="meta")
111+ else:
112+ # 给一个空的合法张量,不能是 nullptr
113+ softmax_lse = torch.empty([], dtype=torch.float32, device="meta")
114+ return (attn_out, softmax_lse)
115+ 
116+# Instantiate the builder
117+sparse_flash_mla_op_builder = SparseFlashMlaOpBuilder()
118+op_module = sparse_flash_mla_op_builder.load() # Compiles/loads the .so file
119+ 
120+ 
121+@impl(AS_LIBRARY, SMLA_METADATA_OP_NAME, "PrivateUse1")
122+def sparse_flash_mla_metadata(
123+ num_heads_q: int, num_heads_kv: int, head_dim: int, cu_seqlens_q: Optional[torch.Tensor] = None,
124+ cu_seqlens_ori_kv: Optional[torch.Tensor] = None, cu_seqlens_cmp_kv: Optional[torch.Tensor] = None,
125+ seqused_q: Optional[torch.Tensor] = None, seqused_ori_kv: Optional[torch.Tensor] = None,
126+ seqused_cmp_kv: Optional[torch.Tensor] = None, cmp_residual_kv: Optional[torch.Tensor] = None,
127+ ori_topk_length: Optional[torch.Tensor] = None, cmp_topk_length: Optional[torch.Tensor] = None,
128+ batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_ori_kv: Optional[int] = None,
129+ max_seqlen_cmp_kv: Optional[int] = None, ori_topk: Optional[int] = None, cmp_topk: Optional[int] = None,
130+ cmp_ratio: Optional[int] = None, ori_mask_mode: Optional[int] = None, cmp_mask_mode: Optional[int] = None,
131+ ori_win_left: Optional[int] = None, ori_win_right: Optional[int] = None, layout_q: Optional[str] = None,
132+ layout_kv: Optional[str] = None, has_ori_kv: Optional[bool] = None, has_cmp_kv: Optional[bool] = None):
133+ """
134+ Dispatcher implementation: NPU.
135+ 'PrivateUse1' is dispatch key for custom NPU backends.
136+ """
137+ batch_size = 0 if batch_size is None else batch_size
138+ max_seqlen_q = 0 if max_seqlen_q is None else max_seqlen_q
139+ max_seqlen_ori_kv = 0 if max_seqlen_ori_kv is None else max_seqlen_ori_kv
140+ max_seqlen_cmp_kv = 0 if max_seqlen_cmp_kv is None else max_seqlen_cmp_kv
141+ ori_topk = 0 if ori_topk is None else ori_topk
142+ cmp_topk = 0 if cmp_topk is None else cmp_topk
143+ cmp_ratio = 1 if cmp_ratio is None else cmp_ratio
144+ ori_mask_mode = 0 if ori_mask_mode is None else ori_mask_mode
145+ cmp_mask_mode = 0 if cmp_mask_mode is None else cmp_mask_mode
146+ ori_win_left = -1 if ori_win_left is None else ori_win_left
147+ ori_win_right = -1 if ori_win_right is None else ori_win_right
148+ layout_q = "BSND" if layout_q is None else layout_q
149+ layout_kv = "BSND" if layout_kv is None else layout_kv
150+ has_ori_kv = True if has_ori_kv is None else has_ori_kv
151+ has_cmp_kv = True if has_cmp_kv is None else has_cmp_kv
152+ 
153+ return op_module.sparse_flash_mla_metadata(
154+ num_heads_q, num_heads_kv, head_dim, cu_seqlens_q, cu_seqlens_ori_kv, cu_seqlens_cmp_kv, seqused_q,
155+ seqused_ori_kv, seqused_cmp_kv, cmp_residual_kv, ori_topk_length, cmp_topk_length, batch_size, max_seqlen_q,
156+ max_seqlen_ori_kv, max_seqlen_cmp_kv, ori_topk, cmp_topk, cmp_ratio, ori_mask_mode, cmp_mask_mode, ori_win_left,
157+ ori_win_right, layout_q, layout_kv, has_ori_kv, has_cmp_kv)
158+ 
159+ 
160+@torch.library.register_kernel("cann_ops_transformer::" + SMLA_METADATA_OP_NAME, None)
161+def sparse_flash_mla_metadata_fallback(
162+ num_heads_q: int, num_heads_kv: int, head_dim: int, cu_seqlens_q: Optional[torch.Tensor] = None,
163+ cu_seqlens_ori_kv: Optional[torch.Tensor] = None, cu_seqlens_cmp_kv: Optional[torch.Tensor] = None,
164+ seqused_q: Optional[torch.Tensor] = None, seqused_ori_kv: Optional[torch.Tensor] = None,
165+ seqused_cmp_kv: Optional[torch.Tensor] = None, cmp_residual_kv: Optional[torch.Tensor] = None,
166+ ori_topk_length: Optional[torch.Tensor] = None, cmp_topk_length: Optional[torch.Tensor] = None,
167+ batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_ori_kv: Optional[int] = None,
168+ max_seqlen_cmp_kv: Optional[int] = None, ori_topk: Optional[int] = None, cmp_topk: Optional[int] = None,
169+ cmp_ratio: Optional[int] = None, ori_mask_mode: Optional[int] = None, cmp_mask_mode: Optional[int] = None,
170+ ori_win_left: Optional[int] = None, ori_win_right: Optional[int] = None, layout_q: Optional[str] = None,
171+ layout_kv: Optional[str] = None, has_ori_kv: Optional[bool] = None, has_cmp_kv: Optional[bool] = None):
172+ # 处理所有 tensor 都为 None 的情况
173+ # 调用 NPU 实现
174+ return sparse_flash_mla_metadata(
175+ num_heads_q, num_heads_kv, head_dim, cu_seqlens_q, cu_seqlens_ori_kv, cu_seqlens_cmp_kv, seqused_q,
176+ seqused_ori_kv, seqused_cmp_kv, cmp_residual_kv, ori_topk_length, cmp_topk_length, batch_size, max_seqlen_q,
177+ max_seqlen_ori_kv, max_seqlen_cmp_kv, ori_topk, cmp_topk, cmp_ratio, ori_mask_mode, cmp_mask_mode, ori_win_left,
178+ ori_win_right, layout_q, layout_kv, has_ori_kv, has_cmp_kv)
179+ 
180+torch.compiler.allow_in_graph(sparse_flash_mla_metadata)
181+ 
182+ 
183+@impl(AS_LIBRARY, sparse_flash_mla_op_builder.name, "PrivateUse1")
184+def sparse_flash_mla(q,
185+ ori_kv=None, cmp_kv=None,
186+ ori_sparse_indices=None, cmp_sparse_indices=None,
187+ ori_block_table=None, cmp_block_table=None,
188+ cu_seqlens_q=None, cu_seqlens_ori_kv=None,
189+ cu_seqlens_cmp_kv=None, seqused_q=None,
190+ seqused_ori_kv=None, seqused_cmp_kv=None,
191+ cmp_residual_kv=None,
192+ ori_topk_length=None, cmp_topk_length=None,
193+ sinks=None, metadata=None,
194+ softmax_scale=1.0, cmp_ratio=1,
195+ ori_mask_mode=4, cmp_mask_mode=3,
196+ ori_win_left=127, ori_win_right=0,
197+ layout_q='BSND', layout_kv='PA_BBND',
198+ topk_value_mode=1, return_softmax_lse=False):
199+ """
200+ dispatcher implementation for NPU.
201+ 'PrivateUse1' is the combine key for custom NPU backends.
202+ """
203+ return op_module.sparse_flash_mla(q,
204+ ori_kv, cmp_kv,
205+ ori_sparse_indices, cmp_sparse_indices,
206+ ori_block_table, cmp_block_table,
207+ cu_seqlens_q, cu_seqlens_ori_kv,
208+ cu_seqlens_cmp_kv, seqused_q,
209+ seqused_ori_kv, seqused_cmp_kv,
210+ cmp_residual_kv,
211+ ori_topk_length, cmp_topk_length,
212+ sinks, metadata,
213+ softmax_scale, cmp_ratio,
214+ ori_mask_mode, cmp_mask_mode,
215+ ori_win_left, ori_win_right,
216+ layout_q, layout_kv,
217+ topk_value_mode, return_softmax_lse)
Rtorch_extension/npu_ops_transformer/ops/sparse_flash_mla_grad.pytorch_extension/cann_ops_transformer/ops/sparse_flash_mla_grad.py+2-2
@@ -10,8 +10,8 @@
10import torch10import torch
11import torch_npu11import torch_npu
12from torch.library import impl12from torch.library import impl
13-from npu_ops_transformer.op_builder.builder import OpBuilder13+from cann_ops_transformer.op_builder.builder import OpBuilder
14-from npu_ops_transformer.op_builder.builder import AS_LIBRARY14+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
15 15 
16 16 
17class SparseFlashMlaGradOpBuilder(OpBuilder):17class SparseFlashMlaGradOpBuilder(OpBuilder):
Rtorch_extension/npu_ops_transformer/ops/sparse_flash_mla_grad_metadata.pytorch_extension/cann_ops_transformer/ops/sparse_flash_mla_grad_metadata.py+2-2
@@ -9,8 +9,8 @@
9import torch9import torch
10import torch_npu10import torch_npu
11from torch.library import impl11from torch.library import impl
12-from npu_ops_transformer.op_builder.builder import AS_LIBRARY12+from cann_ops_transformer.op_builder.builder import AS_LIBRARY
13-from npu_ops_transformer.op_builder.builder import OpBuilder13+from cann_ops_transformer.op_builder.builder import OpBuilder
14 14 
15class SparseFlashMlaGradMetadataOpBuilder(OpBuilder):15class SparseFlashMlaGradMetadataOpBuilder(OpBuilder):
16 def __init__(self):16 def __init__(self):
@@ -1,75 +0,0 @@
1-/**
2- * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3- * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4- * CANN Open Software License Agreement Version 2.0 (the "License").
5- * Please refer to the License for details. You may not use this file except in compliance with the License.
6- * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7- * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8- * See LICENSE in the root of the software repository for the full text of the License.
9- */
10- 
11-/*!
12- * \file lightning_indexer_v2_metadata.cpp
13- * \brief
14- */
15- 
16-#include <torch/extension.h>
17-#include "aclnn_common.h"
18- 
19-namespace op_api {
20-using npu_utils = at_npu::native::NpuUtils;
21-constexpr int64_t OUTPUT_SIZE = 1024;
22- 
23-const c10::optional<at::Tensor> li_v2_get_valid_tensor(const c10::optional<at::Tensor> &tensor_opt, at::Device device)
24-{
25- return tensor_opt.has_value() ? tensor_opt : torch::empty({0}, torch::dtype(torch::kInt32).device(device));
26-};
27- 
28-at::Tensor npu_lightning_indexer_v2_metadata(
29- int64_t num_heads_q, int64_t num_heads_k, int64_t head_dim, int64_t topk,
30- const c10::optional<at::Tensor> &cu_seqlens_q, const c10::optional<at::Tensor> &cu_seqlens_k,
31- const c10::optional<at::Tensor> &seqused_q, const c10::optional<at::Tensor> &seqused_k,
32- const c10::optional<at::Tensor> &cmp_residual_k,
33- int64_t batch_size, int64_t max_seqlen_q, int64_t max_seqlen_k, c10::string_view layout_q,
34- c10::string_view layout_k, int64_t mask_mode, int64_t cmp_ratio)
35-{
36- const c10::string_view device = "npu";
37- at::Device output_device = at::Device(std::string(device));
38- if (cu_seqlens_q.has_value()) {
39- output_device = cu_seqlens_q.value().device();
40- } else if (cu_seqlens_k.has_value()) {
41- output_device = cu_seqlens_k.value().device();
42- } else if (seqused_q.has_value()) {
43- output_device = seqused_q.value().device();
44- } else if (seqused_k.has_value()) {
45- output_device = seqused_k.value().device();
46- } else if (cmp_residual_k.has_value()) {
47- output_device = cmp_residual_k.value().device();
48- }
49- 
50- at::Tensor output = torch::empty({OUTPUT_SIZE}, torch::dtype(torch::kInt32).device(output_device));
51- auto cu_seqlens_q_val = li_v2_get_valid_tensor(cu_seqlens_q, output_device);
52- auto cu_seqlens_k_val = li_v2_get_valid_tensor(cu_seqlens_k, output_device);
53- auto seqused_q_val = li_v2_get_valid_tensor(seqused_q, output_device);
54- auto seqused_k_val = li_v2_get_valid_tensor(seqused_k, output_device);
55- auto cmp_residual_k_val = li_v2_get_valid_tensor(cmp_residual_k, output_device);
56- 
57- // convert str
58- std::string layout_q_str = std::string(layout_q);
59- std::string layout_k_str = std::string(layout_k);
60- char *layout_q_ptr = const_cast<char *>(layout_q_str.c_str());
61- char *layout_k_ptr = const_cast<char *>(layout_k_str.c_str());
62- 
63- ACLNN_CMD(aclnnLightningIndexerV2Metadata, cu_seqlens_q_val, cu_seqlens_k_val, seqused_q_val, seqused_k_val,
64- cmp_residual_k_val,
65- num_heads_q, num_heads_k, head_dim, topk,
66- batch_size, max_seqlen_q, max_seqlen_k, layout_q_ptr, layout_k_ptr, mask_mode, cmp_ratio,
67- output);
68- return output;
69-}
70- 
71-PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
72-{
73- m.def("npu_lightning_indexer_v2_metadata", &npu_lightning_indexer_v2_metadata, "npu_lightning_indexer_v2_metadata");
74-}
75-} // namespace op_api
@@ -1,77 +0,0 @@
1-/**
2- * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3- * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4- * CANN Open Software License Agreement Version 2.0 (the "License").
5- * Please refer to the License for details. You may not use this file except in compliance with the License.
6- * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7- * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8- * See LICENSE in the root of the software repository for the full text of the License.
9- */
10- 
11-/*!
12- * \file quant_lightning_indexer_v2_metadata.cpp
13- * \brief
14- */
15- 
16-#include <torch/extension.h>
17-#include "aclnn_common.h"
18- 
19-namespace op_api {
20-using npu_utils = at_npu::native::NpuUtils;
21-constexpr int64_t OUTPUT_SIZE = 1024;
22- 
23-const c10::optional<at::Tensor> qli_v2_get_valid_tensor(const c10::optional<at::Tensor> &tensor_opt, at::Device device)
24-{
25- return tensor_opt.has_value() ? tensor_opt : torch::empty({0}, torch::dtype(torch::kInt32).device(device));
26-};
27- 
28-at::Tensor npu_quant_lightning_indexer_v2_metadata(
29- int64_t num_heads_q, int64_t num_heads_k, int64_t head_dim, int64_t topk, int64_t q_quant_mode,
30- int64_t k_quant_mode,
31- const c10::optional<at::Tensor> &cu_seqlens_q, const c10::optional<at::Tensor> &cu_seqlens_k,
32- const c10::optional<at::Tensor> &seqused_q, const c10::optional<at::Tensor> &seqused_k,
33- const c10::optional<at::Tensor> &cmp_residual_k,
34- int64_t batch_size, int64_t max_seqlen_q, int64_t max_seqlen_k, c10::string_view layout_q,
35- c10::string_view layout_k, int64_t mask_mode, int64_t cmp_ratio)
36-{
37- const c10::string_view device = "npu";
38- at::Device output_device = at::Device(std::string(device));
39- if (cu_seqlens_q.has_value()) {
40- output_device = cu_seqlens_q.value().device();
41- } else if (cu_seqlens_k.has_value()) {
42- output_device = cu_seqlens_k.value().device();
43- } else if (seqused_q.has_value()) {
44- output_device = seqused_q.value().device();
45- } else if (seqused_k.has_value()) {
46- output_device = seqused_k.value().device();
47- } else if (cmp_residual_k.has_value()) {
48- output_device = cmp_residual_k.value().device();
49- }
50- 
51- at::Tensor output = torch::empty({OUTPUT_SIZE}, torch::dtype(torch::kInt32).device(output_device));
52- auto cu_seqlens_q_val = qli_v2_get_valid_tensor(cu_seqlens_q, output_device);
53- auto cu_seqlens_k_val = qli_v2_get_valid_tensor(cu_seqlens_k, output_device);
54- auto seqused_q_val = qli_v2_get_valid_tensor(seqused_q, output_device);
55- auto seqused_k_val = qli_v2_get_valid_tensor(seqused_k, output_device);
56- auto cmp_residual_k_val = qli_v2_get_valid_tensor(cmp_residual_k, output_device);
57- 
58- // convert str
59- std::string layout_q_str = std::string(layout_q);
60- std::string layout_k_str = std::string(layout_k);
61- char *layout_q_ptr = const_cast<char *>(layout_q_str.c_str());
62- char *layout_k_ptr = const_cast<char *>(layout_k_str.c_str());
63- 
64- ACLNN_CMD(aclnnQuantLightningIndexerV2Metadata, cu_seqlens_q_val, cu_seqlens_k_val, seqused_q_val, seqused_k_val,
65- cmp_residual_k_val,
66- num_heads_q, num_heads_k, head_dim, topk, q_quant_mode, k_quant_mode,
67- batch_size, max_seqlen_q, max_seqlen_k, layout_q_ptr, layout_k_ptr, mask_mode, cmp_ratio,
68- output);
69- return output;
70-}
71- 
72-PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
73-{
74- m.def("npu_quant_lightning_indexer_v2_metadata", &npu_quant_lightning_indexer_v2_metadata,
75- "npu_quant_lightning_indexer_v2_metadata");
76-}
77-} // namespace op_api
@@ -1,98 +0,0 @@
1-# Copyright (c) 2026 Huawei Technologies Co., Ltd.
2-# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
3-# CANN Open Software License Agreement Version 2.0 (the "License").
4-# Please refer to the License for details. You may not use this file except in compliance with the License.
5-# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
6-# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
7-# See LICENSE in the root of the software repository for the full text of the License.
8- 
9-from typing import Optional
10-import torch
11-import torchair
12-from torch.library import impl
13-from npu_ops_transformer.op_builder.builder import OpBuilder
14-from npu_ops_transformer.op_builder.builder import AS_LIBRARY
15-from torchair.ge._ge_graph import Tensor, TensorSpec, DataType
16- 
17- 
18-class LightningIndexerV2MetadataOpBuilder(OpBuilder):
19- def __init__(self):
20- super(LightningIndexerV2MetadataOpBuilder, self).__init__("npu_lightning_indexer_v2_metadata")
21- 
22- def sources(self):
23- """Path to C++ source code."""
24- return ['ops/csrc/lightning_indexer_v2_metadata.cpp']
25- 
26- def schema(self) -> str:
27- """PyTorch operator signature."""
28- return "npu_lightning_indexer_v2_metadata(int num_heads_q, int num_heads_k, int head_dim, int topk, *, " \
29- "Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_k=None, Tensor? seqused_q=None, Tensor? seqused_k=None, " \
30- "Tensor? cmp_residual_k=None, int batch_size=0, int max_seqlen_q=0, int max_seqlen_k=0, " \
31- "str layout_q='BSND', str layout_k='BSND', int mask_mode=0, int cmp_ratio=1) -> Tensor"
32- 
33- def register_meta(self):
34- """
35- Registers Meta implementation (Shape/Dtype inference).
36- Essential for Autograd and FakeTensor support.
37- """
38- @torch.library.register_fake("npu_ops_transformer::" + self.name)
39- def npu_lightning_indexer_v2_metadata_meta(
40- num_heads_q: int, num_heads_k: int, head_dim: int, topk: int,
41- cu_seqlens_q: Optional[Tensor] = None, cu_seqlens_k: Optional[Tensor] = None,
42- seqused_q: Optional[Tensor] = None, seqused_k: Optional[Tensor] = None,
43- cmp_residual_k: Optional[Tensor] = None,
44- batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_k: Optional[int] = None,
45- layout_q: Optional[str] = None, layout_k: Optional[str] = None, mask_mode: Optional[int] = None,
46- cmp_ratio: Optional[int] = None):
47- return torch.empty((1024), dtype=torch.int32, device="npu")
48- 
49-# Instantiate the builder
50-lightning_indexer_v2_metadata_op_builder = LightningIndexerV2MetadataOpBuilder()
51-op_module = lightning_indexer_v2_metadata_op_builder.load()
52- 
53- 
54-@impl(AS_LIBRARY, lightning_indexer_v2_metadata_op_builder.name, "PrivateUse1")
55-def npu_lightning_indexer_v2_metadata(
56- num_heads_q: int, num_heads_k: int, head_dim: int, topk: int,
57- cu_seqlens_q: Optional[Tensor] = None, cu_seqlens_k: Optional[Tensor] = None, seqused_q: Optional[Tensor] = None,
58- seqused_k: Optional[Tensor] = None, cmp_residual_k: Optional[Tensor] = None,
59- batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_k: Optional[int] = None,
60- layout_q: Optional[str] = None, layout_k: Optional[str] = None, mask_mode: Optional[int] = None,
61- cmp_ratio: Optional[int] = None):
62- """
63- Dispatcher implementation: NPU.
64- 'PrivateUse1' is dispatch key for custom NPU backends.
65- """
66- batch_size = 0 if batch_size is None else batch_size
67- max_seqlen_q = 0 if max_seqlen_q is None else max_seqlen_q
68- max_seqlen_k = 0 if max_seqlen_k is None else max_seqlen_k
69- layout_q = "BSND" if layout_q is None else layout_q
70- layout_k = "BSND" if layout_k is None else layout_k
71- mask_mode = 0 if mask_mode is None else mask_mode
72- cmp_ratio = 1 if cmp_ratio is None else cmp_ratio
73- 
74- return op_module.npu_lightning_indexer_v2_metadata(
75- num_heads_q, num_heads_k, head_dim, topk,
76- cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k, cmp_residual_k,
77- batch_size, max_seqlen_q, max_seqlen_k, layout_q, layout_k, mask_mode, cmp_ratio)
78- 
79- 
80-@torch.library.register_kernel("npu_ops_transformer::npu_lightning_indexer_v2_metadata", None)
81-def npu_lightning_indexer_v2_metadata_fallback(
82- num_heads_q: int, num_heads_k: int, head_dim: int, topk: int,
83- cu_seqlens_q: Optional[Tensor] = None, cu_seqlens_k: Optional[Tensor] = None, seqused_q: Optional[Tensor] = None,
84- seqused_k: Optional[Tensor] = None, cmp_residual_k: Optional[Tensor] = None,
85- batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_k: Optional[int] = None,
86- layout_q: Optional[str] = None, layout_k: Optional[str] = None, mask_mode: Optional[int] = None,
87- cmp_ratio: Optional[int] = None):
88- # 处理所有 tensor 都为 None 的情况
89- # 可以在这里创建一个 NPU tensor 来触发 PrivateUse1 backend
90- if all(t is None for t in [cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k, cmp_residual_k]):
91- _ = torch.empty(1, dtype=torch.int32, device="npu")
92- # 调用 NPU 实现
93- return npu_lightning_indexer_v2_metadata(
94- num_heads_q, num_heads_k, head_dim, topk,
95- cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k, cmp_residual_k,
96- batch_size, max_seqlen_q, max_seqlen_k, layout_q, layout_k, mask_mode, cmp_ratio)
97-
98-torch.compiler.allow_in_graph(npu_lightning_indexer_v2_metadata)
@@ -1,99 +0,0 @@
1-# Copyright (c) 2026 Huawei Technologies Co., Ltd.
2-# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
3-# CANN Open Software License Agreement Version 2.0 (the "License").
4-# Please refer to the License for details. You may not use this file except in compliance with the License.
5-# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
6-# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
7-# See LICENSE in the root of the software repository for the full text of the License.
8- 
9-from typing import Optional
10-import torch
11-import torchair
12-from torch.library import impl
13-from npu_ops_transformer.op_builder.builder import OpBuilder
14-from npu_ops_transformer.op_builder.builder import AS_LIBRARY
15-from torchair.ge._ge_graph import Tensor, TensorSpec, DataType
16- 
17- 
18-class QuantLightningIndexerV2MetadataOpBuilder(OpBuilder):
19- def __init__(self):
20- super(QuantLightningIndexerV2MetadataOpBuilder, self).__init__("npu_quant_lightning_indexer_v2_metadata")
21- 
22- def sources(self):
23- """Path to C++ source code."""
24- return ['ops/csrc/quant_lightning_indexer_v2_metadata.cpp']
25- 
26- def schema(self) -> str:
27- """PyTorch operator signature."""
28- return "npu_quant_lightning_indexer_v2_metadata(int num_heads_q, int num_heads_k, int head_dim, int topk, " \
29- "int q_quant_mode, int k_quant_mode, *, Tensor? cu_seqlens_q=None, Tensor? cu_seqlens_k=None, " \
30- "Tensor? seqused_q=None, Tensor? seqused_k=None, Tensor? cmp_residual_k=None, int batch_size=0, " \
31- "int max_seqlen_q=0, int max_seqlen_k=0, str layout_q='BSND', str layout_k='BSND', int mask_mode=0, " \
32- "int cmp_ratio=1) -> Tensor"
33- 
34- def register_meta(self):
35- """
36- Registers Meta implementation (Shape/Dtype inference).
37- Essential for Autograd and FakeTensor support.
38- """
39- @torch.library.register_fake("npu_ops_transformer::" + self.name)
40- def npu_quant_lightning_indexer_v2_metadata_meta(
41- num_heads_q: int, num_heads_k: int, head_dim: int, topk: int, q_quant_mode: int, k_quant_mode: int,
42- cu_seqlens_q: Optional[Tensor] = None, cu_seqlens_k: Optional[Tensor] = None,
43- seqused_q: Optional[Tensor] = None, seqused_k: Optional[Tensor] = None,
44- cmp_residual_k: Optional[Tensor] = None,
45- batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_k: Optional[int] = None,
46- layout_q: Optional[str] = None, layout_k: Optional[str] = None, mask_mode: Optional[int] = None,
47- cmp_ratio: Optional[int] = None):
48- return torch.empty((1024), dtype=torch.int32, device="npu")
49- 
50-# Instantiate the builder
51-quant_lightning_indexer_v2_metadata_op_builder = QuantLightningIndexerV2MetadataOpBuilder()
52-op_module = quant_lightning_indexer_v2_metadata_op_builder.load()
53- 
54- 
55-@impl(AS_LIBRARY, quant_lightning_indexer_v2_metadata_op_builder.name, "PrivateUse1")
56-def npu_quant_lightning_indexer_v2_metadata(
57- num_heads_q: int, num_heads_k: int, head_dim: int, topk: int, q_quant_mode: int, k_quant_mode: int,
58- cu_seqlens_q: Optional[Tensor] = None, cu_seqlens_k: Optional[Tensor] = None, seqused_q: Optional[Tensor] = None,
59- seqused_k: Optional[Tensor] = None, cmp_residual_k: Optional[Tensor] = None,
60- batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_k: Optional[int] = None,
61- layout_q: Optional[str] = None, layout_k: Optional[str] = None, mask_mode: Optional[int] = None,
62- cmp_ratio: Optional[int] = None):
63- """
64- Dispatcher implementation: NPU.
65- 'PrivateUse1' is dispatch key for custom NPU backends.
66- """
67- batch_size = 0 if batch_size is None else batch_size
68- max_seqlen_q = 0 if max_seqlen_q is None else max_seqlen_q
69- max_seqlen_k = 0 if max_seqlen_k is None else max_seqlen_k
70- layout_q = "BSND" if layout_q is None else layout_q
71- layout_k = "BSND" if layout_k is None else layout_k
72- mask_mode = 0 if mask_mode is None else mask_mode
73- cmp_ratio = 1 if cmp_ratio is None else cmp_ratio
74-
75- return op_module.npu_quant_lightning_indexer_v2_metadata(
76- num_heads_q, num_heads_k, head_dim, topk, q_quant_mode, k_quant_mode,
77- cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k, cmp_residual_k,
78- batch_size, max_seqlen_q, max_seqlen_k, layout_q, layout_k, mask_mode, cmp_ratio)
79- 
80- 
81-@torch.library.register_kernel("npu_ops_transformer::npu_quant_lightning_indexer_v2_metadata", None)
82-def npu_quant_lightning_indexer_v2_metadata_fallback(
83- num_heads_q: int, num_heads_k: int, head_dim: int, topk: int, q_quant_mode: int, k_quant_mode: int,
84- cu_seqlens_q: Optional[Tensor] = None, cu_seqlens_k: Optional[Tensor] = None, seqused_q: Optional[Tensor] = None,
85- seqused_k: Optional[Tensor] = None, cmp_residual_k: Optional[Tensor] = None,
86- batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_k: Optional[int] = None,
87- layout_q: Optional[str] = None, layout_k: Optional[str] = None, mask_mode: Optional[int] = None,
88- cmp_ratio: Optional[int] = None):
89- # 处理所有 tensor 都为 None 的情况
90- # 可以在这里创建一个 NPU tensor 来触发 PrivateUse1 backend
91- if all(t is None for t in [cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k, cmp_residual_k]):
92- _ = torch.empty(1, dtype=torch.int32, device="npu")
93- # 调用 NPU 实现
94- return npu_quant_lightning_indexer_v2_metadata(
95- num_heads_q, num_heads_k, head_dim, topk, q_quant_mode, k_quant_mode,
96- cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k, cmp_residual_k,
97- batch_size, max_seqlen_q, max_seqlen_k, layout_q, layout_k, mask_mode, cmp_ratio)
98-
99-torch.compiler.allow_in_graph(npu_quant_lightning_indexer_v2_metadata)
@@ -9,9 +9,9 @@
9# -----------------------------------------------------------------------------------------------------------9# -----------------------------------------------------------------------------------------------------------
10 10 
11 11 
12-__description__ = "NpuOpsTransformer"12+__description__ = "CannOpsTransformer"
13__version__ = "1.0.0"13__version__ = "1.0.0"
14-__package_name__ = 'npu_ops_transformer'14+__package_name__ = 'cann_ops_transformer'
15 15 
16import os16import os
17import sys17import sys