已合并
torch_extension目录整改 #6973
jinying创建于 6月16日
torch_extension目录整改 #6973
已合并
共 85 个文件变更+1756-1205
| @@ -2,8 +2,8 @@ import torch | |||
| 2 | import torch_npu | 2 | import torch_npu |
| 3 | import math | 3 | import math |
| 4 | import numpy as np | 4 | import numpy as np |
| 5 | -import npu_ops_transformer | 5 | +import cann_ops_transformer |
| 6 | -from npu_ops_transformer.ops import npu_flash_attn | 6 | +from cann_ops_transformer.ops import npu_flash_attn |
| 7 | torch.manual_seed(42) | 7 | torch.manual_seed(42) |
| 8 | 8 | ||
| 9 | # B = 1 | 9 | # B = 1 |
| @@ -15,9 +15,9 @@ import math | |||
| 15 | import numpy as np | 15 | import numpy as np |
| 16 | import random | 16 | import random |
| 17 | from einops import rearrange | 17 | from einops import rearrange |
| 18 | -import npu_ops_transformer | 18 | +import cann_ops_transformer |
| 19 | -from npu_ops_transformer.ops import npu_flash_attn | 19 | +from cann_ops_transformer.ops import npu_flash_attn |
| 20 | -from npu_ops_transformer.ops import npu_flash_attn_metadata | 20 | +from cann_ops_transformer.ops import npu_flash_attn_metadata |
| 21 | from utils import trans_bnsd_to_layout | 21 | from utils import trans_bnsd_to_layout |
| 22 | import torchair | 22 | import torchair |
| 23 | from torchair.configs.compiler_config import CompilerConfig | 23 | from 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 | |||
| 8 | from .base import Backend | 8 | from .base import Backend |
| 9 | 9 | ||
| 10 | try: | 10 | try: |
| 11 | - from npu_ops_transformer.ops import npu_flash_attn, npu_flash_attn_metadata | 11 | + from cann_ops_transformer.ops import npu_flash_attn, npu_flash_attn_metadata |
| 12 | _HAS_NPU = True | 12 | _HAS_NPU = True |
| 13 | except ImportError: | 13 | except ImportError: |
| 14 | _HAS_NPU = False | 14 | _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 | ```bash | 113 | ```bash |
| 114 | # NPU | 114 | # 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 | # GPU | 117 | # GPU |
| 118 | python -c "import torch, flash_attn, einops; print('GPU OK')" | 118 | python -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 | ```bash | 431 | ```bash |
| 432 | python test_flash_attn.py --case_id BASE_01 --use_gpu | 432 | python test_flash_attn.py --case_id BASE_01 --use_gpu |
| 433 | ``` | 433 | ``` |
| @@ -1,7 +1,7 @@ | |||
| 1 | import subprocess | 1 | import subprocess |
| 2 | import torch | 2 | import torch |
| 3 | import torch_npu | 3 | import torch_npu |
| 4 | -import npu_ops_transformer | 4 | +import cann_ops_transformer |
| 5 | 5 | ||
| 6 | # 初始化 NPU | 6 | # 初始化 NPU |
| 7 | torch_npu.npu.set_device(0) | 7 | torch_npu.npu.set_device(0) |
| @@ -40,8 +40,8 @@ print("seqused_kv",seqused_kv) | |||
| 40 | print("seqused_kv",seqused_kv) | 40 | print("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 | |||
| 19 | import math | 19 | import math |
| 20 | import ctypes | 20 | import ctypes |
| 21 | import copy | 21 | import copy |
| 22 | -import npu_ops_transformer | 22 | +import cann_ops_transformer |
| 23 | -from npu_ops_transformer.ops import npu_lightning_indexer_v2 | 23 | +from cann_ops_transformer.ops import lightning_indexer_v2 |
| 24 | 24 | ||
| 25 | class GeneralizedLIV2: | 25 | class 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 = None | 630 | 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.py | 83 | + 详见torch_extension/cann_ops_transformer/ops/mixed_quant_sparse_flash_mla.py |
Mattention/mixed_quant_sparse_flash_mla/tests/pytest/batch/mixed_quant_sparse_flash_mla_process.py+8-8
| @@ -18,7 +18,7 @@ import pytest | |||
| 18 | import random | 18 | import random |
| 19 | import numpy as np | 19 | import numpy as np |
| 20 | import math | 20 | import math |
| 21 | -import npu_ops_transformer | 21 | +import cann_ops_transformer |
| 22 | import custom_ops as ops | 22 | import custom_ops as ops |
| 23 | import torchair | 23 | import torchair |
| 24 | from torchair.configs.compiler_config import CompilerConfig | 24 | from 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 | |||
| 19 | import math | 19 | import math |
| 20 | import ctypes | 20 | import ctypes |
| 21 | import copy | 21 | import copy |
| 22 | -import npu_ops_transformer | 22 | +import cann_ops_transformer |
| 23 | 23 | ||
| 24 | FP32_FRACTION_BITS = 23 # fp32尾数位数 | 24 | FP32_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 | |||
| 19 | import torch | 19 | import torch |
| 20 | import torch_npu | 20 | import torch_npu |
| 21 | 21 | ||
| 22 | -# Register npu_sparse_flash_mla and npu_quant_lightning_indexer_v2_metadata via PTA | 22 | +# Register sparse_flash_mla and sparse_flash_mla_metadata via PTA |
| 23 | TORCH_EXT_PATH = os.path.normpath(os.path.join(os.path.dirname(os.path.abspath(__file__)), | 23 | TORCH_EXT_PATH = os.path.normpath(os.path.join(os.path.dirname(os.path.abspath(__file__)), |
| 24 | '../../../../torch_extension')) | 24 | '../../../../torch_extension')) |
| 25 | if TORCH_EXT_PATH not in sys.path: | 25 | if 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: F401 | 27 | +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: F401 | 28 | +from cann_ops_transformer.ops import sparse_flash_mla_metadata as _metadata_registration # noqa: F401 |
O | |||
| 29 | 29 | ||
| 30 | class Network(torch.nn.Module): | 30 | class Network(torch.nn.Module): |
O 【openlibing.ci】识别到代码检查告警抑制注释,匹配工具:flake8,请Committer检视其合理性。 ![]() ![]() | |||
| 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 | # 生成 metadata | 91 | # 生成 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_lse | 155 | 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" | |||
| 33 | if MINDSPEED_PATH.exists(): | 33 | if 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-import | 36 | +import cann_ops_transformer # noqa: E402,F401 pylint: disable=wrong-import-position,unused-import |
O 【openlibing.ci】识别到代码检查告警抑制注释,匹配工具:flake8,请Committer检视其合理性。 ![]() ![]() | |||
| 37 | 37 | ||
| 38 | 38 | ||
| 39 | SLI_METADATA_SIZE = 64 | 39 | SLI_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/script | 127 | 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/whl | 135 | 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}" ]; then | 380 | 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 | fi | 404 | fi |
| 405 | 405 | ||
| 406 | - logandprint "[INFO]: npu_ops_transformer whl package installed successfully" | 406 | + logandprint "[INFO]: cann_ops_transformer whl package installed successfully" |
| 407 | else | 407 | else |
| 408 | - logandprint "[INFO]: No npu_ops_transformer whl package found, skipping" | 408 | + logandprint "[INFO]: No cann_ops_transformer whl package found, skipping" |
| 409 | fi | 409 | fi |
| 410 | } | 410 | } |
| 411 | 411 | ||
| @@ -183,19 +183,19 @@ remove_init_py() { | |||
| 183 | remove_whl_package() { | 183 | remove_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" ]; then | 186 | + 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; then | 190 | if command -v pip3 &>/dev/null; then |
| 191 | - pip3 uninstall -y npu_ops_transformer --target="${python_dir}" 2>/dev/null | 191 | + pip3 uninstall -y cann_ops_transformer --target="${python_dir}" 2>/dev/null |
| 192 | fi | 192 | fi |
| 193 | 193 | ||
| 194 | # 直接删除目录(确保清理干净) | 194 | # 直接删除目录(确保清理干净) |
| 195 | - rm -rf "${python_dir}/npu_ops_transformer" 2>/dev/null | 195 | + rm -rf "${python_dir}/cann_ops_transformer" 2>/dev/null |
| 196 | - rm -f ${python_dir}/npu_ops_transformer-*.whl 2>/dev/null | 196 | + rm -f ${python_dir}/cann_ops_transformer-*.whl 2>/dev/null |
| 197 | - rm -rf ${python_dir}/npu_ops_transformer-*.egg-info 2>/dev/null | 197 | + rm -rf ${python_dir}/cann_ops_transformer-*.egg-info 2>/dev/null |
| 198 | - rm -rf ${python_dir}/npu_ops_transformer-*.dist-info 2>/dev/null | 198 | + 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)" ]; then | 201 | if [ -d "${python_dir}" ] && [ -z "$(ls -A ${python_dir} 2>/dev/null)" ]; then |
| @@ -207,7 +207,7 @@ remove_whl_package() { | |||
| 207 | fi | 207 | fi |
| 208 | fi | 208 | fi |
| 209 | 209 | ||
| 210 | - logandprint "[INFO]: npu_ops_transformer whl package removed" | 210 | + logandprint "[INFO]: cann_ops_transformer whl package removed" |
| 211 | fi | 211 | fi |
| 212 | } | 212 | } |
| 213 | 213 | ||
| @@ -1,6 +1,6 @@ | |||
| 1 | -# NPU Ops Transformer | 1 | +# 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 & Installation | 5 | ## Build & Installation |
| 6 | 6 | ||
| @@ -37,19 +37,19 @@ | |||
| 37 | 37 | ||
| 38 | ## Quick Start | 38 | ## 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 | ```python | 42 | ```python |
| 43 | import torch | 43 | import torch |
| 44 | import torch_npu | 44 | import torch_npu |
| 45 | -import npu_ops_transformer | 45 | +import cann_ops_transformer |
| 46 | 46 | ||
| 47 | # Initialize data on NPU | 47 | # Initialize data on NPU |
| 48 | x = torch.randn(10, 32, dtype=torch.float32).npu() | 48 | x = torch.randn(10, 32, dtype=torch.float32).npu() |
| 49 | 49 | ||
| 50 | # Call the custom NPU operator | 50 | # Call the custom NPU operator |
| 51 | # This triggers JIT compilation on the first call | 51 | # 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 implementation | 54 | # Verify against CPU ATen implementation |
| 55 | cpu_x = x.cpu() | 55 | cpu_x = x.cpu() |
| @@ -102,8 +102,8 @@ This file manages the JIT compilation logic and registers the operator into the | |||
| 102 | import torch | 102 | import torch |
| 103 | import torch_npu | 103 | import torch_npu |
| 104 | from torch.library import impl | 104 | from torch.library import impl |
| 105 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 105 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 106 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 106 | +from cann_ops_transformer.op_builder.builder import AS_LIBRARY |
| 107 | 107 | ||
| 108 | class AbsOpBuilder(OpBuilder): | 108 | class AbsOpBuilder(OpBuilder): |
| 109 | def __init__(self): | 109 | def __init__(self): |
Rtorch_extension/npu_ops_transformer/__init__.py→torch_extension/cann_ops_transformer/__init__.py+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/common/inc/aclnn_common.h→torch_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.h | 12 | * \file aclnn_common.h |
| 13 | * \brief | 13 | * \brief |
| 14 | */ | 14 | */ |
| 15 | 15 | ||
| 16 | -#ifndef NPU_OPS_TRANSFORMER_ACLNN_COMMON_H | 16 | +#ifndef CANN_OPS_TRANSFORMER_ACLNN_COMMON_H |
| 17 | -#define NPU_OPS_TRANSFORMER_ACLNN_COMMON_H | 17 | +#define CANN_OPS_TRANSFORMER_ACLNN_COMMON_H |
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | 20 | ||
| @@ -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 | + | ||
| 776 | template <typename T> | 781 | template <typename T> |
| 777 | void Release(T value) | 782 | void 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_H | 964 | +#endif // CANN_OPS_TRANSFORMER_ACLNN_COMMON_H |
Rtorch_extension/npu_ops_transformer/common/inc/hccl_common.h→torch_extension/cann_ops_transformer/common/inc/hccl_common.h+7-5
| @@ -13,8 +13,8 @@ | |||
| 13 | * \brief | 13 | * \brief |
| 14 | */ | 14 | */ |
| 15 | 15 | ||
| 16 | -#ifndef NPU_OPS_TRANSFORMER_HCCL_COMMON_H | 16 | +#ifndef CANN_OPS_TRANSFORMER_HCCL_COMMON_H |
| 17 | -#define NPU_OPS_TRANSFORMER_HCCL_COMMON_H | 17 | +#define CANN_OPS_TRANSFORMER_HCCL_COMMON_H |
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | 20 | ||
| @@ -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获取groupHandle | 156 | + 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"); // 创建HcclContext | 159 | 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_H | 198 | +#endif // CANN_OPS_TRANSFORMER_HCCL_COMMON_H |
Rtorch_extension/npu_ops_transformer/doc/get_low_latency_ccl_buffer_size.md→torch_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 | |||
| 76 | import os | 76 | import os |
| 77 | import torch | 77 | import torch |
| 78 | import torch_npu | 78 | import torch_npu |
| 79 | -from npu_ops_transformer.ops import MoeDistributeBuffer | 79 | +from cann_ops_transformer.ops import MoeDistributeBuffer |
| 80 | 80 | ||
| 81 | server_num = 1 | 81 | server_num = 1 |
| 82 | dev_num = 16 | 82 | dev_num = 16 |
Rtorch_extension/npu_ops_transformer/doc/mega_moe.md→torch_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 dist | 181 | import torch.distributed as dist |
| 182 | from torch.distributed import ReduceOp | 182 | from torch.distributed import ReduceOp |
| 183 | import torch.multiprocessing as mp | 183 | import torch.multiprocessing as mp |
| 184 | - from npu_ops_transformer.ops import get_symm_buffer_for_mega_moe, mega_moe | 184 | + from cann_ops_transformer.ops import get_symm_buffer_for_mega_moe, mega_moe |
| 185 | import torchair | 185 | import torchair |
| 186 | 186 | ||
| 187 | E = 4 | 187 | E = 4 |
Rtorch_extension/npu_ops_transformer/doc/mhc_post.md→torch_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 | ```python | 197 | ```python |
| 198 | import torch | 198 | import torch |
| 199 | import torch_npu | 199 | import torch_npu |
| 200 | - from npu_ops_transformer.ops import mhc_post | 200 | + from cann_ops_transformer.ops import mhc_post |
| 201 | 201 | ||
| 202 | B = 2 | 202 | B = 2 |
| 203 | S = 8 | 203 | S = 8 |
Rtorch_extension/npu_ops_transformer/doc/mhc_pre_sinkhorn.md→torch_extension/cann_ops_transformer/doc/mhc_pre_sinkhorn.md+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/doc/npu_flash_attn.md→torch_extension/cann_ops_transformer/doc/npu_flash_attn.md+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/doc/npu_get_mega_moe_ccl_buffer_size.md→torch_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 | |||
| 68 | import os | 68 | import os |
| 69 | import torch | 69 | import torch |
| 70 | import torch_npu | 70 | import torch_npu |
| 71 | -from npu_ops_transformer.ops import npu_get_mega_moe_ccl_buffer_size | 71 | +from cann_ops_transformer.ops import npu_get_mega_moe_ccl_buffer_size |
| 72 | 72 | ||
| 73 | server_num = 1 | 73 | server_num = 1 |
| 74 | rank_per_dev = 2 | 74 | rank_per_dev = 2 |
Rtorch_extension/npu_ops_transformer/doc/npu_low_latency_combine.md→torch_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 Process | 148 | from torch.multiprocessing import Process |
| 149 | import torch.distributed as dist | 149 | import torch.distributed as dist |
| 150 | from torch.distributed import ReduceOp | 150 | from torch.distributed import ReduceOp |
| 151 | - from npu_ops_transformer.ops import MoeDistributeBuffer | 151 | + 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 dist | 352 | import torch.distributed as dist |
| 353 | from torch.distributed import ReduceOp | 353 | from torch.distributed import ReduceOp |
| 354 | import time | 354 | import time |
| 355 | - from npu_ops_transformer.ops import MoeDistributeBuffer | 355 | + 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.md→torch_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 Process | 165 | from torch.multiprocessing import Process |
| 166 | import torch.distributed as dist | 166 | import torch.distributed as dist |
| 167 | from torch.distributed import ReduceOp | 167 | from torch.distributed import ReduceOp |
| 168 | - from npu_ops_transformer.ops import MoeDistributeBuffer | 168 | + 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 dist | 369 | import torch.distributed as dist |
| 370 | from torch.distributed import ReduceOp | 370 | from torch.distributed import ReduceOp |
| 371 | import time | 371 | import time |
| 372 | - from npu_ops_transformer.ops import MoeDistributeBuffer | 372 | + 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.py→torch_extension/cann_ops_transformer/op_builder/builder.py+3-3
| @@ -15,10 +15,10 @@ import torch | |||
| 15 | from torch.utils.cpp_extension import load | 15 | from torch.utils.cpp_extension import load |
| 16 | from torch.library import Library | 16 | from torch.library import Library |
| 17 | import torch_npu | 17 | import torch_npu |
| 18 | -import npu_ops_transformer | 18 | +import cann_ops_transformer |
| 19 | 19 | ||
| 20 | ASCEND_HOME_PATH = "ASCEND_HOME_PATH" | 20 | ASCEND_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 | ||
| 24 | class OpBuilder(ABC): | 24 | class OpBuilder(ABC): |
| @@ -34,7 +34,7 @@ class OpBuilder(ABC): | |||
| 34 | self.name = name | 34 | 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__.py→torch_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 of | 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"). | 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 | |||
| 23 | from .flash_attn import npu_flash_attn | 23 | from .flash_attn import npu_flash_attn |
| 24 | from .graph_convert.graph_convert_flash_attn import convert_npu_flash_attn | 24 | from .graph_convert.graph_convert_flash_attn import convert_npu_flash_attn |
| 25 | from .flash_attn_metadata import npu_flash_attn_metadata | 25 | from .flash_attn_metadata import npu_flash_attn_metadata |
| 26 | -from .lightning_indexer_v2_metadata import npu_lightning_indexer_v2_metadata | 26 | +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_metadata | 27 | +from .graph_convert.graph_convert_mixed_quant_sparse_flash_mla import ( |
| 28 | -from .mixed_quant_sparse_flash_mla import npu_mixed_quant_sparse_flash_mla | 28 | + 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 | ||
| 33 | from .npu_sparse_lightning_indexer_kl_loss_grad import npu_sparse_lightning_indexer_kl_loss_grad | 32 | from .npu_sparse_lightning_indexer_kl_loss_grad import npu_sparse_lightning_indexer_kl_loss_grad |
| 34 | from .npu_sparse_lightning_indexer_kl_loss_grad_metadata import npu_sparse_lightning_indexer_kl_loss_grad_metadata | 33 | from .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_v2 | 34 | +from .lightning_indexer_v2 import lightning_indexer_v2, lightning_indexer_metadata |
| 36 | -from .quant_lightning_indexer_v2 import npu_quant_lightning_indexer_v2 | 35 | +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 | +) | ||
| 37 | from .mhc_post import mhc_post | 40 | from .mhc_post import mhc_post |
| 38 | from .mhc_pre_sinkhorn import mhc_pre_sinkhorn | 41 | from .mhc_pre_sinkhorn import mhc_pre_sinkhorn |
Rtorch_extension/npu_ops_transformer/ops/comm_context.py→torch_extension/cann_ops_transformer/ops/comm_context.py+1-1
| @@ -9,7 +9,7 @@ | |||
| 9 | # ----------------------------------------------------------------------------------------------------------- | 9 | # ----------------------------------------------------------------------------------------------------------- |
| 10 | import torch | 10 | import torch |
| 11 | import torch_npu | 11 | import torch_npu |
| 12 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 12 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | class CommContextOpBuilder(OpBuilder): | 15 | class CommContextOpBuilder(OpBuilder): |
Rtorch_extension/npu_ops_transformer/ops/csrc/comm_context.cpp→torch_extension/cann_ops_transformer/ops/csrc/comm_context.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/flash_attn.cpp→torch_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 of | 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"). | 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.cpp | 12 | + * \file npu_flash_attn.cpp |
| 13 | - * \brief | 13 | + * \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, |
已过期 torch_extension 目录整改涉及公共头/宏迁移。请确认所有 OpBuilder 仍指向新路径且 ![]() ![]() | |||
| 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 module | 85 | +// 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_api | 90 | } // namespace op_api |
Rtorch_extension/npu_ops_transformer/ops/csrc/flash_attn_metadata.cpp→torch_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 of | 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"). | 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.cpp | 12 | + * \file flash_attn_metadata.cpp |
| 13 | - * \brief | 13 | + * \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_api | 40 | +} // namespace op_api |
Rtorch_extension/npu_ops_transformer/ops/csrc/lightning_indexer_v2.cpp→torch_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 of | 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"). | 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.cpp | 12 | +* \file lightning_indexer_v2.cpp |
| 13 | -* \brief | 13 | +* \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 size | 22 | +// 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 | -// 工具函数,推导输出shape | 29 | +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 tensor | 87 | + 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_output | 88 | + 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 str | 93 | + 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 module | 102 | + |
| 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_api | 148 | } // namespace op_api |
Rtorch_extension/npu_ops_transformer/ops/csrc/mega_moe.cpp→torch_extension/cann_ops_transformer/ops/csrc/mega_moe.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/mhc_post.cpp→torch_extension/cann_ops_transformer/ops/csrc/mhc_post.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/mhc_post_backward.cpp→torch_extension/cann_ops_transformer/ops/csrc/mhc_post_backward.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/mhc_pre_sinkhorn.cpp→torch_extension/cann_ops_transformer/ops/csrc/mhc_pre_sinkhorn.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/mhc_pre_sinkhorn_backward.cpp→torch_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.cpp→torch_extension/cann_ops_transformer/ops/csrc/mixed_quant_sparse_flash_mla.cpp+64-2
| @@ -25,6 +25,66 @@ const int DIM_2 = 2; | |||
| 25 | const int DIM_3 = 3; | 25 | const int DIM_3 = 3; |
| 26 | const int DIM_4 = 4; | 26 | const 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 | + | ||
| 28 | std::tuple<at::Tensor, at::Tensor> construct_mixed_quant_sparse_flash_mla_atten_out_tensor( | 88 | std::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 | ||
| 159 | PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) | 219 | PYBIND11_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_api | 225 | } // namespace op_api |
Rtorch_extension/npu_ops_transformer/ops/csrc/moe_distribute_combine.cpp→torch_extension/cann_ops_transformer/ops/csrc/moe_distribute_combine.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/csrc/moe_distribute_dispatch.cpp→torch_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.cpp→torch_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.cpp→torch_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.cpp→torch_extension/cann_ops_transformer/ops/csrc/quant_lightning_indexer_v2.cpp+44-2
| @@ -26,6 +26,46 @@ 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 | +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 | // 工具函数,推导输出shape | 69 | // 工具函数,推导输出shape |
| 30 | std::tuple<at::Tensor, at::Tensor> construct_quant_lightning_indexer_output_tensor( | 70 | std::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 module | 142 | // Bind the C++ function to Python module |
| 103 | PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) | 143 | PYBIND11_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_api | 149 | } // 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 | + | ||
| 16 | + | ||
| 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.cpp→torch_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.cpp→torch_extension/cann_ops_transformer/ops/csrc/sparse_flash_mla_grad_metadata.cpp+0-0
文件重命名但无更改。
Rtorch_extension/npu_ops_transformer/ops/deep_ep.py→torch_extension/cann_ops_transformer/ops/deep_ep.py+4-4
| @@ -11,8 +11,8 @@ import torch | |||
| 11 | import torch_npu | 11 | import torch_npu |
| 12 | from torch.library import impl | 12 | from torch.library import impl |
| 13 | from torch_npu.utils._error_code import ErrCode, ops_error | 13 | from torch_npu.utils._error_code import ErrCode, ops_error |
| 14 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 14 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 15 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 15 | +from cann_ops_transformer.op_builder.builder import AS_LIBRARY |
| 16 | from .moe_distribute_combine import npu_moe_distribute_combine | 16 | from .moe_distribute_combine import npu_moe_distribute_combine |
| 17 | from .moe_distribute_dispatch import npu_moe_distribute_dispatch | 17 | from .moe_distribute_dispatch import npu_moe_distribute_dispatch |
| 18 | from .comm_context import CommContextManager | 18 | from .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.py→torch_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 of | 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"). | 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 torch | 10 | +import torch |
| 11 | -import torch_npu | 11 | +import torch_npu |
| 12 | -from torch.library import impl | 12 | +from torch.library import impl |
| 13 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 13 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 14 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 14 | +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 builder | 93 | +# 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 file | 95 | +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.py→torch_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 of | 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"). | 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 torch | 10 | +import torch |
| 11 | -import torch_npu | 11 | +import torch_npu |
| 12 | -from torch.library import impl | 12 | +from torch.library import impl |
| 13 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 13 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 14 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 14 | +from cann_ops_transformer.op_builder.builder import AS_LIBRARY |
| 15 | -from typing import Optional | 15 | +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 batchSize | 25 | + 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 alignedSize | 31 | + 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 builder | 65 | +# 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_size | 82 | + 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_q | 83 | + 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_kv | 84 | + 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_mode | 85 | + mask_mode = 1 if mask_mode is None else mask_mode |
| 86 | - win_left = -1 if win_left is None else win_left | 86 | + win_left = -1 if win_left is None else win_left |
| 87 | - win_right = -1 if win_right is None else win_right | 87 | + win_right = -1 if win_right is None else win_right |
| 88 | - layout_q = "BSND" if layout_q is None else layout_q | 88 | + layout_q = "BSND" if layout_q is None else layout_q |
| 89 | - layout_kv = "BSND" if layout_kv is None else layout_kv | 89 | + layout_kv = "BSND" if layout_kv is None else layout_kv |
| 90 | - layout_out = "BSND" if layout_out is None else layout_out | 90 | + 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.py→torch_extension/cann_ops_transformer/ops/graph_convert/graph_convert_flash_attn.py+101-101
| @@ -1,102 +1,102 @@ | |||
| 1 | -try: | 1 | +try: |
| 2 | - import torch | 2 | + import torch |
| 3 | - import torch_npu | 3 | + import torch_npu |
| 4 | - import torchair | 4 | + import torchair |
| 5 | - from torch.library import impl | 5 | + from torch.library import impl |
| 6 | - from torchair._ge_concrete_graph import ge_apis as ge | 6 | + from torchair._ge_concrete_graph import ge_apis as ge |
| 7 | - from torchair.ge._ge_graph import Tensor, TensorSpec | 7 | + from torchair.ge._ge_graph import Tensor, TensorSpec |
| 8 | - from torchair._ge_concrete_graph.fx2ge_converter import declare_supported, register_fx_node_ge_converter | 8 | + from torchair._ge_concrete_graph.fx2ge_converter import declare_supported, register_fx_node_ge_converter |
| 9 | - from torchair._ge_concrete_graph.supported_declaration import Support | 9 | + from torchair._ge_concrete_graph.supported_declaration import Support |
| 10 | - from typing import Any, Dict, List, Tuple, Union, Callable, Optional | 10 | + from typing import Any, Dict, List, Tuple, Union, Callable, Optional |
| 11 | - from torchair._ge_concrete_graph.ge_ir_pb2 import GraphDef, OpDef, TensorDescriptor, TensorDef | 11 | + 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_name | 12 | + from torchair.ge._ge_graph import get_default_ge_graph, next_unique_name |
| 13 | - from torchair.ge._ge_graph import auto_convert_to_tensor | 13 | + from torchair.ge._ge_graph import auto_convert_to_tensor |
| 14 | - from torchair.ge._ge_graph import Tensor, TensorSpec, DataType, TensorType | 14 | + from torchair.ge._ge_graph import Tensor, TensorSpec, DataType, TensorType |
| 15 | - from torchair.ge._ge_graph import compat_as_bytes, compat_as_bytes_list | 15 | + 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_float | 16 | + 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_desc | 17 | + from torchair.ge._ge_graph import get_invalid_desc |
| 18 | - from torchair._ge_concrete_graph.compat_ir import ge_op, IrDef | 18 | + from torchair._ge_concrete_graph.compat_ir import ge_op, IrDef |
| 19 | - from torchair.ge import attr | 19 | + from torchair.ge import attr |
| 20 | - _TORCHAIR_AVAILABLE = True | 20 | + _TORCHAIR_AVAILABLE = True |
| 21 | -except ImportError: | 21 | +except ImportError: |
| 22 | - _TORCHAIR_AVAILABLE = False | 22 | + _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 result | 52 | + 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.py→torch_extension/cann_ops_transformer/ops/graph_convert/graph_convert_lightning_indexer.py+3-3
| @@ -33,8 +33,8 @@ except ImportError: | |||
| 33 | _TORCHAIR_AVAILABLE = False | 33 | _TORCHAIR_AVAILABLE = False |
| 34 | 34 | ||
| 35 | if _TORCHAIR_AVAILABLE: | 35 | if _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.py→torch_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, |
Atorch_extension/cann_ops_transformer/ops/graph_convert/graph_convert_mixed_quant_sparse_flash_mla.py+49-0
| @@ -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 | + | ||
| 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.py→torch_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 x | 254 | 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.py→torch_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.py→torch_extension/cann_ops_transformer/ops/graph_convert/graph_convert_quant_lightning_indexer.py+4-4
| @@ -33,13 +33,13 @@ except ImportError: | |||
| 33 | _TORCHAIR_AVAILABLE = False | 33 | _TORCHAIR_AVAILABLE = False |
| 34 | 34 | ||
| 35 | if _TORCHAIR_AVAILABLE: | 35 | if _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 | + | ||
| 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.py→torch_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 of | 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"). | 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 torch | 10 | +from typing import Optional |
| 11 | -import torch_npu | 11 | +import torch |
| 12 | -from torch.library import impl | 12 | +import torch_npu |
| 13 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 13 | +from torch.library import impl |
| 14 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 14 | +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 builder | 65 | + 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 file | 67 | + 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.zhe | 77 | + 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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.py→torch_extension/cann_ops_transformer/ops/mega_moe.py+3-3
| @@ -13,8 +13,8 @@ import torch | |||
| 13 | import torch_npu | 13 | import torch_npu |
| 14 | from torch.library import impl | 14 | from torch.library import impl |
| 15 | from torch_npu.utils._error_code import ErrCode, ops_error | 15 | from torch_npu.utils._error_code import ErrCode, ops_error |
| 16 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 16 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 17 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 17 | +from cann_ops_transformer.op_builder.builder import AS_LIBRARY |
| 18 | from .comm_context import CommContextManager | 18 | from .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.py→torch_extension/cann_ops_transformer/ops/mhc_post.py+3-3
| @@ -10,8 +10,8 @@ | |||
| 10 | import torch | 10 | import torch |
| 11 | import torch_npu | 11 | import torch_npu |
| 12 | from torch.library import impl | 12 | from torch.library import impl |
| 13 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 13 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 14 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 14 | +from cann_ops_transformer.op_builder.builder import AS_LIBRARY |
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | class MhcPostFunction(torch.autograd.Function): | 17 | class MhcPostFunction(torch.autograd.Function): |
| @@ -24,7 +24,7 @@ class MhcPostFunction(torch.autograd.Function): | |||
| 24 | 24 | ||
| 25 | def backward(ctx, grad_output): | 25 | def backward(ctx, grad_output): |
| 26 | x, h_res, h_out, h_post = ctx.saved_tensors | 26 | x, h_res, h_out, h_post = ctx.saved_tensors |
| 27 | - from npu_ops_transformer.ops.mhc_post_backward import mhc_post_backward | 27 | + 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_post | 29 | grad_output, x, h_res, h_out, h_post |
| 30 | ) | 30 | ) |
Rtorch_extension/npu_ops_transformer/ops/mhc_post_backward.py→torch_extension/cann_ops_transformer/ops/mhc_post_backward.py+2-2
| @@ -10,8 +10,8 @@ | |||
| 10 | import torch | 10 | import torch |
| 11 | import torch_npu | 11 | import torch_npu |
| 12 | from torch.library import impl | 12 | from torch.library import impl |
| 13 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 13 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 14 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 14 | +from cann_ops_transformer.op_builder.builder import AS_LIBRARY |
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | class MhcPostBackwardOpBuilder(OpBuilder): | 17 | class MhcPostBackwardOpBuilder(OpBuilder): |
Rtorch_extension/npu_ops_transformer/ops/mhc_pre_sinkhorn.py→torch_extension/cann_ops_transformer/ops/mhc_pre_sinkhorn.py+3-3
| @@ -10,8 +10,8 @@ | |||
| 10 | import torch | 10 | import torch |
| 11 | import torch_npu | 11 | import torch_npu |
| 12 | from torch.library import impl | 12 | from torch.library import impl |
| 13 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 13 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 14 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 14 | +from cann_ops_transformer.op_builder.builder import AS_LIBRARY |
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | class MhcPreSinkhornFunction(torch.autograd.Function): | 17 | class 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_tensors | 33 | x, phi, alpha, bias, h_pre, hc_before_norm, inv_rms, sum_out, norm_out = ctx.saved_tensors |
| 34 | hc_eps = ctx.hc_eps | 34 | hc_eps = ctx.hc_eps |
| 35 | 35 | ||
| 36 | - from npu_ops_transformer.ops.mhc_pre_sinkhorn_backward import mhc_pre_sinkhorn_backward | 36 | + 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.py→torch_extension/cann_ops_transformer/ops/mhc_pre_sinkhorn_backward.py+2-2
| @@ -10,8 +10,8 @@ | |||
| 10 | import torch | 10 | import torch |
| 11 | import torch_npu | 11 | import torch_npu |
| 12 | from torch.library import impl | 12 | from torch.library import impl |
| 13 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 13 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 14 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 14 | +from cann_ops_transformer.op_builder.builder import AS_LIBRARY |
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | class MhcPreSinkhornBackwardOpBuilder(OpBuilder): | 17 | class MhcPreSinkhornBackwardOpBuilder(OpBuilder): |
Rtorch_extension/npu_ops_transformer/ops/mixed_quant_sparse_flash_mla.py→torch_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 | ||
| 9 | import torch | 10 | import torch |
| 10 | import torch_npu | 11 | import torch_npu |
| 11 | from torch.library import impl | 12 | from torch.library import impl |
| 12 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 13 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 13 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 14 | +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 | ||
| 16 | class MixedQuantSparseFlashMlaOpBuilder(OpBuilder): | 19 | class 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 | + | ||
| 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 | 82 | ||
| 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 builder | 119 | # 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 file | 121 | +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 | + | ||
| 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 | + | ||
| 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 NPU | 208 | 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.py→torch_extension/cann_ops_transformer/ops/moe_distribute_combine.py+2-2
| @@ -10,8 +10,8 @@ | |||
| 10 | import torch | 10 | import torch |
| 11 | import torch_npu | 11 | import torch_npu |
| 12 | from torch.library import impl | 12 | from torch.library import impl |
| 13 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 13 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 14 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 14 | +from cann_ops_transformer.op_builder.builder import AS_LIBRARY |
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | class MoeDistributeCombineOpBuilder(OpBuilder): | 17 | class MoeDistributeCombineOpBuilder(OpBuilder): |
Rtorch_extension/npu_ops_transformer/ops/moe_distribute_dispatch.py→torch_extension/cann_ops_transformer/ops/moe_distribute_dispatch.py+2-2
| @@ -12,8 +12,8 @@ import torch | |||
| 12 | import torch_npu | 12 | import torch_npu |
| 13 | from torch.library import impl | 13 | from torch.library import impl |
| 14 | from torch_npu.utils._error_code import ErrCode, ops_error | 14 | from torch_npu.utils._error_code import ErrCode, ops_error |
| 15 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 15 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 16 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 16 | +from cann_ops_transformer.op_builder.builder import AS_LIBRARY |
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | TORCH_DTYPE_ENUM_VALUE_TO_SCALAR_TYPE_MAP = { | 19 | TORCH_DTYPE_ENUM_VALUE_TO_SCALAR_TYPE_MAP = { |
Rtorch_extension/npu_ops_transformer/ops/npu_sparse_lightning_indexer_kl_loss_grad.py→torch_extension/cann_ops_transformer/ops/npu_sparse_lightning_indexer_kl_loss_grad.py+3-3
| @@ -12,8 +12,8 @@ import torch | |||
| 12 | import torch_npu | 12 | import torch_npu |
| 13 | from torch.library import impl | 13 | from torch.library import impl |
| 14 | 14 | ||
| 15 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 15 | +from cann_ops_transformer.op_builder.builder import AS_LIBRARY |
| 16 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 16 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | class SparseLightningIndexerKLLossGradOpBuilder(OpBuilder): | 19 | class SparseLightningIndexerKLLossGradOpBuilder(OpBuilder): |
| @@ -117,7 +117,7 @@ except ImportError: | |||
| 117 | 117 | ||
| 118 | if _TORCHAIR_AVAILABLE: | 118 | if _TORCHAIR_AVAILABLE: |
| 119 | 119 | ||
| 120 | - torch.ops.npu_ops_transformer.npu_sparse_lightning_indexer_kl_loss_grad.default | 120 | + 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.py→torch_extension/cann_ops_transformer/ops/npu_sparse_lightning_indexer_kl_loss_grad_metadata.py+4-4
| @@ -12,8 +12,8 @@ import torch | |||
| 12 | import torch_npu | 12 | import torch_npu |
| 13 | from torch.library import impl | 13 | from torch.library import impl |
| 14 | 14 | ||
| 15 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 15 | +from cann_ops_transformer.op_builder.builder import AS_LIBRARY |
| 16 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 16 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | SLI_KL_LOSS_GRAD_METADATA_SIZE = 64 | 19 | SLI_KL_LOSS_GRAD_METADATA_SIZE = 64 |
| @@ -118,7 +118,7 @@ def npu_sparse_lightning_indexer_kl_loss_grad_metadata( | |||
| 118 | 118 | ||
| 119 | 119 | ||
| 120 | 120 | ||
| 121 | - "npu_ops_transformer::npu_sparse_lightning_indexer_kl_loss_grad_metadata", None | 121 | + "cann_ops_transformer::npu_sparse_lightning_indexer_kl_loss_grad_metadata", None |
| 122 | ) | 122 | ) |
| 123 | def npu_sparse_lightning_indexer_kl_loss_grad_metadata_fallback( | 123 | def 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 | ||
| 173 | if _TORCHAIR_AVAILABLE: | 173 | if _TORCHAIR_AVAILABLE: |
| 174 | 174 | ||
| 175 | - torch.ops.npu_ops_transformer.npu_sparse_lightning_indexer_kl_loss_grad_metadata.default | 175 | + 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.py→torch_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 | ||
| 10 | import torch | 11 | import torch |
| 11 | import torch_npu | 12 | import torch_npu |
| 12 | from torch.library import impl | 13 | from torch.library import impl |
| 13 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 14 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 14 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 15 | +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 | ||
| 17 | class QuantLightningIndexerV2OpBuilder(OpBuilder): | 20 | class 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 | + | ||
| 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 | 61 | ||
| 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 builder | 90 | # 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 file | 92 | +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 | + | ||
| 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 | + | ||
| 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.zhe | 145 | 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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.py→torch_extension/cann_ops_transformer/ops/sparse_flash_mla_grad.py+2-2
| @@ -10,8 +10,8 @@ | |||
| 10 | import torch | 10 | import torch |
| 11 | import torch_npu | 11 | import torch_npu |
| 12 | from torch.library import impl | 12 | from torch.library import impl |
| 13 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 13 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 14 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 14 | +from cann_ops_transformer.op_builder.builder import AS_LIBRARY |
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | class SparseFlashMlaGradOpBuilder(OpBuilder): | 17 | class SparseFlashMlaGradOpBuilder(OpBuilder): |
Rtorch_extension/npu_ops_transformer/ops/sparse_flash_mla_grad_metadata.py→torch_extension/cann_ops_transformer/ops/sparse_flash_mla_grad_metadata.py+2-2
| @@ -9,8 +9,8 @@ | |||
| 9 | import torch | 9 | import torch |
| 10 | import torch_npu | 10 | import torch_npu |
| 11 | from torch.library import impl | 11 | from torch.library import impl |
| 12 | -from npu_ops_transformer.op_builder.builder import AS_LIBRARY | 12 | +from cann_ops_transformer.op_builder.builder import AS_LIBRARY |
| 13 | -from npu_ops_transformer.op_builder.builder import OpBuilder | 13 | +from cann_ops_transformer.op_builder.builder import OpBuilder |
| 14 | 14 | ||
| 15 | class SparseFlashMlaGradMetadataOpBuilder(OpBuilder): | 15 | class 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 | - | ||
| 17 | - | ||
| 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 | - | ||
| 17 | - | ||
| 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 | - | ||
| 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 | - | ||
| 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 | - | ||
| 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 | - | ||
| 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 | - | ||
| 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 | - | ||
| 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 | ||
| 16 | import os | 16 | import os |
| 17 | import sys | 17 | import sys |


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