已合并
[feat] 新增triton_dense_to_jagged算子atk测试迁移(npu部分) #123
QianZH97创建于 27 天前
[feat] 新增triton_dense_to_jagged算子atk测试迁移(npu部分) #123
已合并
共 6 个文件变更+966-0
| @@ -0,0 +1,156 @@ | |||
| 1 | +# triton_dense_to_jagged | ||
| 2 | + | ||
| 3 | +## 来源 | ||
| 4 | + | ||
| 5 | +`triton_dense_to_jagged` 算子源自 FBGEMM 代码仓: | ||
| 6 | + | ||
| 7 | +<https://github.com/pytorch/FBGEMM/blob/v1.5.0-release/fbgemm_gpu/fbgemm_gpu/triton/jagged/triton_jagged_tensor_ops.py> | ||
| 8 | + | ||
| 9 | +## 算子功能 | ||
| 10 | + | ||
| 11 | +`triton_dense_to_jagged` 用于将 **dense tensor(稠密张量)** 转换为 **jagged tensor(不规则张量)**,并支持在写出时与一个 jagged tensor 做元素级算术融合(`add` 或 `mul`)。 | ||
| 12 | + | ||
| 13 | +- 通过 `jagged_offsets` 描述 dense 中哪些切片被抽取、抽取后的 jagged tensor 形状如何。 | ||
| 14 | +- 对于多维 jagged(`JAGGED_DIM > 2`),通过 `dense_indices` 把 `pid` 映射到 dense 张量中的偏移;若对应位置越界,`dense_indices` 标记为 `-1`,kernel 跳过该 row。 | ||
| 15 | +- 可选地将 dense 与同形状的 `operation_jagged_values` 做元素级融合,结果直接写回 `output_jagged_values`,无需中间临时张量。 | ||
| 16 | + | ||
| 17 | +--- | ||
| 18 | + | ||
| 19 | +## 输入 | ||
| 20 | + | ||
| 21 | +`triton_dense_to_jagged` 算子接受以下输入: | ||
| 22 | + | ||
| 23 | +| 参数 | 说明 | | ||
| 24 | +|------|------| | ||
| 25 | +| `jagged_value_ptr` | 输出 jagged tensor 的扁平值指针,shape `(N_total, inner_dim)`,`dtype` 支持 fp32/fp16/bf16 | | ||
| 26 | +| `jagged_offsets_ptr` | 最内层分段偏移,要求首元素为 `0`、末元素为 `N_total` | | ||
| 27 | +| `jagged_value_row_stride` | 输出 jagged tensor 在行维度上的 stride(通常等于 `inner_dim`) | | ||
| 28 | +| `output_dense_ptr` | 输入 dense 张量指针;`JAGGED_DIM == 2` 时按 `pid * dense_matrix_stride` 偏移,`JAGGED_DIM > 2` 时按 `dense_indices[pid]` 偏移 | | ||
| 29 | +| `dense_indices_ptr` | 多维 jagged 时使用的 dense 索引;1D / 2D jagged 传任意张量即可,越界位置填 `-1` | | ||
| 30 | +| `dense_col_stride` | dense 最后一维 stride | | ||
| 31 | +| `dense_row_stride` | dense 倒数第二维 stride | | ||
| 32 | +| `dense_matrix_stride` | dense 倒数第三维 stride | | ||
| 33 | +| `JAGGED_DIM` | jagged tensor 的维度(`constexpr`),等于 `len(jagged_offsets) + 1`;1D / 2D 时无需 `dense_indices` | | ||
| 34 | +| `thread_block_row_size` | 块行大小(`constexpr`,NPU 端会通过 `triton.autotune` 自动选择;GPU 端常用 `32`) | | ||
| 35 | +| `thread_block_col_size` | 块列大小(`constexpr`,NPU 端会通过 `triton.autotune` 自动选择;GPU 端常用 `32`) | | ||
| 36 | +| `operation_function` | 融合操作,`"add"` / `"mul"` / `None` | | ||
| 37 | +| `operation_jagged_value_ptr` | 形状与 `output_jagged_value` 一致,提供融合操作的右操作数;不融合时传 `None` | | ||
| 38 | + | ||
| 39 | +## 输出 | ||
| 40 | + | ||
| 41 | +`output_jagged_value`:形状 `(N_total, inner_dim)`,其中 `N_total = jagged_offsets[-1][-1]`,`inner_dim = dense.size(-1)`,`dtype` 与 `dense` 相同。对于越界位置(`dense_indices[pid] == -1`)的 row,kernel 不会从 dense 读取,但若 `operation_function` 不为 `None`,仍会进行融合计算(值会被 mask 填 0 后参与)。 | ||
| 42 | + | ||
| 43 | +--- | ||
| 44 | + | ||
| 45 | +## 使用样例 | ||
| 46 | + | ||
| 47 | +### 示例 :直接调用 kernel | ||
| 48 | + | ||
| 49 | +```python | ||
| 50 | +import torch | ||
| 51 | +import torch_npu | ||
| 52 | +from triton_dense_to_jagged import triton_dense_to_jagged | ||
| 53 | + | ||
| 54 | +DEVICE = "npu:0" | ||
| 55 | +dense = torch.arange(1 * 5 * 3, dtype=torch.float32, device=DEVICE).reshape(1, 5, 3) | ||
| 56 | +# jagged_offsets: shape [B+1],每个 batch 抽取 seq 个元素(这里只取前 3 个,裁掉后 2 个) | ||
| 57 | +jagged_offsets = torch.tensor([0, 3], dtype=torch.int64, device=DEVICE) | ||
| 58 | +# 输出 jagged tensor 的 shape 为 (N_total, D) = (3, 3) | ||
| 59 | +output_jagged_value = torch.empty((3, 3), dtype=torch.float32, device=DEVICE) | ||
| 60 | +dense_indices = torch.tensor([0, 1, 2], dtype=torch.int32, device=DEVICE) | ||
| 61 | + | ||
| 62 | +grid = (jagged_offsets.size(0) - 1,) | ||
| 63 | +triton_dense_to_jagged[grid]( | ||
| 64 | + output_jagged_value, # 输出:jagged tensor | ||
| 65 | + jagged_offsets, # 输入:offsets | ||
| 66 | + output_jagged_value.stride(0), # 输出 jagged tensor 行 stride | ||
| 67 | + dense, # 输入:dense tensor | ||
| 68 | + dense_indices, # dense indices | ||
| 69 | + dense.stride(-1), # dense_col_stride | ||
| 70 | + dense.stride(-2), # dense_row_stride | ||
| 71 | + dense.stride(-3), # dense_matrix_stride | ||
| 72 | + JAGGED_DIM=2, # 1 层 offsets -> len(offsets) + 1 = 2 | ||
| 73 | + operation_function=None, | ||
| 74 | + operation_jagged_value_ptr=None, | ||
| 75 | +) | ||
| 76 | + | ||
| 77 | +torch.npu.synchronize(DEVICE) | ||
| 78 | +kernel_out = output_jagged_value.cpu() | ||
| 79 | +print(dense) | ||
| 80 | +# output_jagged_value: [[0.0, 1.0, 2.0], [3.0, 4.0, 5.0], [6.0, 7.0, 8.0]] | ||
| 81 | +print("output_jagged_value:", kernel_out.tolist()) | ||
| 82 | + | ||
| 83 | +``` | ||
| 84 | + | ||
| 85 | +--- | ||
| 86 | + | ||
| 87 | +## 关键约束 | ||
| 88 | + | ||
| 89 | +- `jagged_offsets[-1][0] == 0` 且 `jagged_offsets[-1][-1] == N_total`。 | ||
| 90 | +- `jagged_offsets[i]` 单调递增;相邻两层 offsets 满足 `jagged_offsets[i][-1] == jagged_offsets[i+1].size(0) - 1`。 | ||
| 91 | +- `JAGGED_DIM == 2` 时 dense 至少为 3D;`JAGGED_DIM > 2` 时 dense 至少为 4D,且 `dense_indices[pid]` 越界时填 `-1`。 | ||
| 92 | +- `operation_jagged_values` 的 `dtype` 与 `dense` 保持一致;当 `operation_function` 为 `None` 时 `operation_jagged_value_ptr` 传 `None`。 | ||
| 93 | +- dense 最后一维(特征维 `D`)建议 `<= 8192`,且内部维满足 `B < seq < D`(3D)或 `B1 < B2 < seq < D`(4D)的算子使用约束。 | ||
| 94 | + | ||
| 95 | +--- | ||
| 96 | + | ||
| 97 | +## 目录文件 | ||
| 98 | + | ||
| 99 | +```text | ||
| 100 | +triton_dense_to_jagged/ | ||
| 101 | +├── README.md | ||
| 102 | +├── gpu_atk_test/ | ||
| 103 | +│ ├── triton_dense_to_jagged.py # 原版算子 kernel 实现 | ||
| 104 | +│ └── triton_dense_to_jagged_api.py # GPU 端 ATK 自定义 API 封装 | ||
| 105 | +└── npu_atk_test/ | ||
| 106 | + ├── triton_dense_to_jagged.py # NPU 适配优化后算子 kernel 实现(带 autotune) | ||
| 107 | + ├── triton_dense_to_jagged_api.py # NPU 端 ATK 自定义 API 封装 | ||
| 108 | + ├── triton_dense_to_jagged.yaml # ATK 算子用例参数 | ||
| 109 | + ├── generate_triton_dense_to_jagged.py # 限制参数生成脚本 | ||
| 110 | + └── triton_dense_to_jagged_performance.json # 性能测试样例数据 | ||
| 111 | +``` | ||
| 112 | + | ||
| 113 | +--- | ||
| 114 | + | ||
| 115 | +## ATK 测试 | ||
| 116 | + | ||
| 117 | +测试在 NPU / GPU 两端协同进行: | ||
| 118 | + | ||
| 119 | +- `npu_atk_test/` 部署到 NPU 端;`gpu_atk_test/` 部署到 GPU 端。 | ||
| 120 | +- 精度测试用例在 NPU 端生成。 | ||
| 121 | +- 性能测试用例复用 `npu_atk_test/triton_dense_to_jagged_performance.json`。 | ||
| 122 | +- GPU 端启动 ATK server,NPU 端作为 ATK 客户端拉起精度/性能测试。 | ||
| 123 | + | ||
| 124 | +### 生成精度测试数据 | ||
| 125 | + | ||
| 126 | +```bash | ||
| 127 | +atk case -f triton_dense_to_jagged.yaml -p generate_triton_dense_to_jagged.py | ||
| 128 | +``` | ||
| 129 | + | ||
| 130 | +生成的精度测试用例位于 `result/triton_dense_to_jagged/json/` 下,文件名为 `all_triton_dense_to_jagged.json`。 | ||
| 131 | + | ||
| 132 | +### 远端 GPU 服务器起监听命令 | ||
| 133 | + | ||
| 134 | +```bash | ||
| 135 | +atk server --devices 0 --plugin_path /path/to/api文件所在文件夹 | ||
| 136 | +``` | ||
| 137 | + | ||
| 138 | +### 精度测试 | ||
| 139 | + | ||
| 140 | +```bash | ||
| 141 | +atk node --backend npu --devices 0 node --backend gpu -h {gpu端ip} -p {gpu端监听端口} --devices 0 \ | ||
| 142 | + task -c all_triton_dense_to_jagged.json --task accuracy \ | ||
| 143 | + -p triton_dense_to_jagged_api.py | ||
| 144 | +``` | ||
| 145 | + | ||
| 146 | +注:如需保存性能对比数据则需要在指令后加 --save_data profile | ||
| 147 | + | ||
| 148 | +### 性能测试 | ||
| 149 | + | ||
| 150 | +```bash | ||
| 151 | +atk node --backend npu --devices 0 node --backend gpu -h {gpu端ip} -p {gpu端监听端口} --devices 0 \ | ||
| 152 | + task -c triton_dense_to_jagged_performance.json --task performance_device \ | ||
| 153 | + -p triton_dense_to_jagged_api.py | ||
| 154 | +``` | ||
| 155 | + | ||
| 156 | +注:如需保存性能对比数据则需要在指令后加 --save_data profile | ||
Aexperimental/triton/atk_test/triton_dense_to_jagged/npu_atk_test/generate_triton_dense_to_jagged.py+170-0
| @@ -0,0 +1,170 @@ | |||
| 1 | +# Copyright (c) Huawei Technologies Co., Ltd. 2026. All rights reserved. | ||
| 2 | +# | ||
| 3 | +# Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 4 | +# you may not use this file except in compliance with the License. | ||
| 5 | +# You may obtain a copy of the License at | ||
| 6 | +# | ||
| 7 | +# http://www.apache.org/licenses/LICENSE-2.0 | ||
| 8 | +# | ||
| 9 | +# Unless required by applicable law or agreed to in writing, software | ||
| 10 | +# distributed under the License is distributed on an "AS IS" BASIS, | ||
| 11 | +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| 12 | +# See the License for the specific language governing permissions and | ||
| 13 | +# limitations under the License. | ||
| 14 | +# ============================================================================== | ||
| 15 | +import copy | ||
| 16 | + | ||
| 17 | +from atk.case_generator.generator.generate_types import GENERATOR_REGISTRY | ||
| 18 | +from atk.case_generator.generator.base_generator import CaseGenerator | ||
| 19 | +from atk.configs.case_config import CaseConfig | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | +class DenseToJaggedGenerator(CaseGenerator): | ||
| 24 | + """dense_to_jagged 算子参数约束生成器 | ||
| 25 | + | ||
| 26 | + 在用例生成后,对参数进行约束检查和修正。 | ||
| 27 | + """ | ||
| 28 | + | ||
| 29 | + def after_case_config(self, case_config: CaseConfig) -> CaseConfig: | ||
| 30 | + # 获取输入参数(按 yaml 中定义的顺序) | ||
| 31 | + dense = case_config.inputs[0] | ||
| 32 | + jagged_offsets_input = case_config.inputs[1] | ||
| 33 | + operation_function = case_config.inputs[2] if len(case_config.inputs) > 2 else None | ||
| 34 | + operation_jagged_value = case_config.inputs[3] if len(case_config.inputs) > 3 else None | ||
| 35 | + | ||
| 36 | + # 获取 dense tensor 的形状信息 | ||
| 37 | + dense_shape = list(dense.shape) if hasattr(dense, 'shape') else [1, 1, 1] | ||
| 38 | + dense_dim = len(dense_shape) # 维度数量(3 或 4) | ||
| 39 | + D = dense_shape[-1] | ||
| 40 | + | ||
| 41 | + # ==================== 约束1:特征维度 D <= 8192 ==================== | ||
| 42 | + # triton 的 arange 要求 D 是 2 的幂次方,且有上限 | ||
| 43 | + if D > 8192: | ||
| 44 | + D = 8192 | ||
| 45 | + dense_shape[-1] = D | ||
| 46 | + dense.shape = dense_shape | ||
| 47 | + | ||
| 48 | + # ==================== 情况1: 3D dense [B, seq, D] ==================== | ||
| 49 | + # 例如:[2, 4, 8] 表示 B=2, seq=4, D=8 | ||
| 50 | + # 约束:B < seq < D,且 B <= D/4 | ||
| 51 | + if dense_dim == 3: | ||
| 52 | + B = dense_shape[0] # 批次大小 | ||
| 53 | + seq = dense_shape[1] # 序列长度 | ||
| 54 | + | ||
| 55 | + # 确保 B <= D/4(留出足够空间给 seq) | ||
| 56 | + if B > D // 4: | ||
| 57 | + B = max(2, D // 4) | ||
| 58 | + dense_shape[0] = B | ||
| 59 | + | ||
| 60 | + # 确保 B < seq < D | ||
| 61 | + if B >= seq or seq >= D: | ||
| 62 | + seq = (B + D) // 2 # 保证 B < seq < D | ||
| 63 | + dense_shape[1] = seq | ||
| 64 | + dense.shape = dense_shape | ||
| 65 | + | ||
| 66 | + # 更新 jagged_offsets 的约束 | ||
| 67 | + # 3D dense 对应 1 个 offset,shape 为 [B+1] | ||
| 68 | + needed_offsets_count = 1 | ||
| 69 | + if jagged_offsets_input is not None and hasattr(jagged_offsets_input, '__len__'): | ||
| 70 | + # 截断多余的 offset(如果 yaml 生成了多个) | ||
| 71 | + jagged_offsets_input[:] = jagged_offsets_input[:needed_offsets_count] | ||
| 72 | + | ||
| 73 | + if len(jagged_offsets_input) > 0: | ||
| 74 | + offset_tensor = jagged_offsets_input[0] | ||
| 75 | + if hasattr(offset_tensor, 'shape'): | ||
| 76 | + offset_tensor.shape = [B + 1] # shape 必须为 [B+1] | ||
| 77 | + | ||
| 78 | + # ==================== 情况2: 4D dense [B1, B2, seq, D] ==================== | ||
| 79 | + # 例如:[2, 3, 4, 8] 表示 B1=2, B2=3, seq=4, D=8 | ||
| 80 | + # 约束:B1 < B2 < seq < D,且 B1 <= D/4 | ||
| 81 | + elif dense_dim == 4: | ||
| 82 | + B1 = dense_shape[0] # 第一维批次大小 | ||
| 83 | + B2 = dense_shape[1] # 第二维批次大小 | ||
| 84 | + seq = dense_shape[2] # 序列长度 | ||
| 85 | + | ||
| 86 | + # 确保 B1 <= D/4(留出足够空间给 seq) | ||
| 87 | + if B1 > D // 4: | ||
| 88 | + B1 = max(2, D // 4) | ||
| 89 | + dense_shape[0] = B1 | ||
| 90 | + | ||
| 91 | + # 确保 B1 < B2 < seq < D | ||
| 92 | + # 如果 B1 >= B2,调整 B2 | ||
| 93 | + if B1 >= B2: | ||
| 94 | + B2 = B1 + 1 | ||
| 95 | + dense_shape[1] = B2 | ||
| 96 | + | ||
| 97 | + # 如果 B2 >= seq 或 seq >= D,调整 seq(保证 B2 < seq < D) | ||
| 98 | + if B2 >= seq or seq >= D: | ||
| 99 | + seq = (B2 + D) // 2 | ||
| 100 | + dense_shape[2] = seq | ||
| 101 | + dense.shape = dense_shape | ||
| 102 | + | ||
| 103 | + # 更新 jagged_offsets 的约束 | ||
| 104 | + # 4D dense 对应 2 个 offset | ||
| 105 | + # jagged_offsets[0] shape 为 [B1+1] | ||
| 106 | + # jagged_offsets[1] shape 为 [B1*B2+1] | ||
| 107 | + needed_offsets_count = 2 | ||
| 108 | + if jagged_offsets_input is not None and hasattr(jagged_offsets_input, '__len__'): | ||
| 109 | + # 截断多余的 offset,保证层数不超过 2 | ||
| 110 | + jagged_offsets_input[:] = jagged_offsets_input[:needed_offsets_count] | ||
| 111 | + | ||
| 112 | + if len(jagged_offsets_input) == 0: | ||
| 113 | + # 极少见:yaml 未生成任何 offset 层时,无法仅靠 shape 声明补齐, | ||
| 114 | + # 这里不猜测 dtype/数值,保持为空,交由运行时 _generate_jagged_offsets 兜底。 | ||
| 115 | + pass | ||
| 116 | + else: | ||
| 117 | + offset0_tensor = jagged_offsets_input[0] | ||
| 118 | + if hasattr(offset0_tensor, 'shape'): | ||
| 119 | + offset0_tensor.shape = [B1 + 1] | ||
| 120 | + | ||
| 121 | + if len(jagged_offsets_input) < needed_offsets_count: | ||
| 122 | + # yaml 随机出的 offset 层数不足(如 dense 4D 却只有 1 层 offset)。 | ||
| 123 | + # 这里只声明第 2 层的 shape,具体数值由 ATK 按 range_values 生成, | ||
| 124 | + # 运行时再由 api 的 _generate_jagged_offsets 按 dense 维度重建。 | ||
| 125 | + second_level = copy.deepcopy(offset0_tensor) | ||
| 126 | + second_level.shape = [B1 * B2 + 1] | ||
| 127 | + jagged_offsets_input.append(second_level) | ||
| 128 | + else: | ||
| 129 | + jagged_offsets_input[1].shape = [B1 * B2 + 1] | ||
| 130 | + | ||
| 131 | + # ==================== 约束2:jagged_offsets 数据类型一致 ==================== | ||
| 132 | + # 所有 offset tensor 必须使用相同的数据类型(int32 或 int64) | ||
| 133 | + dtypes = set() | ||
| 134 | + first_dtype = None | ||
| 135 | + for offset_tensor in jagged_offsets_input: | ||
| 136 | + if offset_tensor is not None and hasattr(offset_tensor, 'dtype'): | ||
| 137 | + dtypes.add(offset_tensor.dtype) | ||
| 138 | + if first_dtype is None: | ||
| 139 | + first_dtype = offset_tensor.dtype | ||
| 140 | + | ||
| 141 | + if len(dtypes) > 1: | ||
| 142 | + # 多 dtype 时统一为第一个 offset 的 dtype(保留 yaml 的原始选择, | ||
| 143 | + # 避免硬编码为 int64 导致 int32 case 被强行升格) | ||
| 144 | + unified_dtype = first_dtype | ||
| 145 | + for offset_tensor in jagged_offsets_input: | ||
| 146 | + if offset_tensor is not None and hasattr(offset_tensor, 'dtype'): | ||
| 147 | + offset_tensor.dtype = unified_dtype | ||
| 148 | + | ||
| 149 | + # ==================== 约束3:operation_jagged_value 的数据类型 ==================== | ||
| 150 | + # 当 operation_function 不为 null 时,operation_jagged_value 的数据类型需与 dense 一致 | ||
| 151 | + op_func_val = None | ||
| 152 | + if ( | ||
| 153 | + operation_function is not None | ||
| 154 | + and hasattr(operation_function, 'range_values') | ||
| 155 | + and operation_function.range_values | ||
| 156 | + ): | ||
| 157 | + op_func_val = operation_function.range_values[0] | ||
| 158 | + | ||
| 159 | + if op_func_val is None or str(op_func_val) == 'null' or str(op_func_val) == 'None': | ||
| 160 | + # operation_function 为 null 时,不需要处理 operation_jagged_value | ||
| 161 | + pass | ||
| 162 | + else: | ||
| 163 | + # operation_function 为 "add" 或 "mul" 时,operation_jagged_value 的数据类型必须与 dense 相同 | ||
| 164 | + dense_dtype = None | ||
| 165 | + if hasattr(dense, 'dtype'): | ||
| 166 | + dense_dtype = dense.dtype | ||
| 167 | + if operation_jagged_value is not None and dense_dtype is not None: | ||
| 168 | + operation_jagged_value.dtype = dense_dtype | ||
| 169 | + | ||
| 170 | + return case_config | ||
| @@ -0,0 +1,191 @@ | |||
| 1 | +#!/usr/bin/env python3 | ||
| 2 | +# -*- coding: utf-8 -*- | ||
| 3 | +# | ||
| 4 | +# Copyright (c) Meta Platforms, Inc. and affiliates. | ||
| 5 | +# All rights reserved. | ||
| 6 | +# | ||
| 7 | +# This source code is licensed under the BSD-style license found in the | ||
| 8 | +# LICENSE file in the root directory of this source tree. | ||
| 9 | +# | ||
| 10 | +# Copyright (c) Huawei Technologies Co., Ltd. 2026. All rights reserved. | ||
| 11 | +# | ||
| 12 | +# Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 13 | +# you may not use this file except in compliance with the License. | ||
| 14 | +# You may obtain a copy of the License at | ||
| 15 | +# | ||
| 16 | +# http://www.apache.org/licenses/LICENSE-2.0 | ||
| 17 | +# | ||
| 18 | +# Unless required by applicable law or agreed to in writing, software | ||
| 19 | +# distributed under the License is distributed on an "AS IS" BASIS, | ||
| 20 | +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| 21 | +# See the License for the specific language governing permissions and | ||
| 22 | +# limitations under the License. | ||
| 23 | +# ============================================================================== | ||
| 24 | + | ||
| 25 | +# pyre-strict | ||
| 26 | + | ||
| 27 | +# pyre-ignore-all-errors[6] | ||
| 28 | + | ||
| 29 | +import torch | ||
| 30 | +import triton # @manual | ||
| 31 | +import triton.language as tl # @manual | ||
| 32 | + | ||
| 33 | + | ||
| 34 | + | ||
| 35 | +# pyre-fixme[3]: Return type must be annotated. | ||
| 36 | +# pyre-fixme[2]: Parameter must be annotated. | ||
| 37 | +def tensor_elementwise_add(x, y): | ||
| 38 | + return x + y | ||
| 39 | + | ||
| 40 | + | ||
| 41 | + | ||
| 42 | +# pyre-fixme[3]: Return type must be annotated. | ||
| 43 | +# pyre-fixme[2]: Parameter must be annotated. | ||
| 44 | +def tensor_elementwise_mul(x, y): | ||
| 45 | + return x * y | ||
| 46 | + | ||
| 47 | + | ||
| 48 | + | ||
| 49 | + configs=[ | ||
| 50 | + triton.Config({"thread_block_row_size": 16, "thread_block_col_size": 16}), | ||
| 51 | + triton.Config({"thread_block_row_size": 32, "thread_block_col_size": 32}), | ||
| 52 | + triton.Config({"thread_block_row_size": 64, "thread_block_col_size": 64}), | ||
| 53 | + triton.Config({"thread_block_row_size": 128, "thread_block_col_size": 128}), | ||
| 54 | + triton.Config({"thread_block_row_size": 256, "thread_block_col_size": 256}), | ||
| 55 | + ], | ||
| 56 | + key=["JAGGED_DIM"], | ||
| 57 | +) | ||
| 58 | +# each kernel will handle the conversion of one jagged tensor offset range from corresponding dense index | ||
| 59 | + | ||
| 60 | +def triton_dense_to_jagged( | ||
| 61 | + # pyre-fixme[2]: Parameter must be annotated. | ||
| 62 | + jagged_value_ptr, | ||
| 63 | + # pyre-fixme[2]: Parameter must be annotated. | ||
| 64 | + jagged_offsets_ptr, | ||
| 65 | + jagged_value_row_stride: int, | ||
| 66 | + # pyre-fixme[2]: Parameter must be annotated. | ||
| 67 | + output_dense_ptr, | ||
| 68 | + # pyre-fixme[2]: Parameter must be annotated. | ||
| 69 | + dense_indices_ptr, | ||
| 70 | + # pyre-fixme[2]: Parameter must be annotated. | ||
| 71 | + dense_col_stride, # stride of output dense with dimension (z,y,x) | ||
| 72 | + dense_row_stride: int, | ||
| 73 | + # pyre-fixme[2]: Parameter must be annotated. | ||
| 74 | + dense_matrix_stride, | ||
| 75 | + JAGGED_DIM: tl.constexpr, # number of dimension of jagged tensor | ||
| 76 | + thread_block_row_size: tl.constexpr, | ||
| 77 | + thread_block_col_size: tl.constexpr, | ||
| 78 | + operation_function: tl.constexpr, # fusion arithmetic opeartion function and it's input dense | ||
| 79 | + # pyre-fixme[2]: Parameter must be annotated. | ||
| 80 | + operation_jagged_value_ptr, | ||
| 81 | +) -> None: | ||
| 82 | + pid = tl.program_id(0) | ||
| 83 | + | ||
| 84 | + begin = tl.load(jagged_offsets_ptr + pid) | ||
| 85 | + end = tl.load(jagged_offsets_ptr + (pid + 1)) | ||
| 86 | + | ||
| 87 | + # size of the current value offset range (M , N) | ||
| 88 | + N = jagged_value_row_stride | ||
| 89 | + M = end - begin | ||
| 90 | + | ||
| 91 | + dense_boundary_col = dense_row_stride | ||
| 92 | + # tl.minimum will change the return type cased compile issue | ||
| 93 | + # in that case use if statement instead | ||
| 94 | + if N < dense_row_stride: | ||
| 95 | + dense_boundary_col = N | ||
| 96 | + | ||
| 97 | + dense_boundary_row = tl.minimum(dense_matrix_stride // dense_row_stride, M) | ||
| 98 | + | ||
| 99 | + jagged_value_ptr += begin * jagged_value_row_stride | ||
| 100 | + if JAGGED_DIM > 2: | ||
| 101 | + dense_indice = tl.load(dense_indices_ptr + pid) | ||
| 102 | + # if dense output range we set dense_boundary to -1 | ||
| 103 | + # that mean dense values will not be use with mask | ||
| 104 | + # since we still need the calculation of fusion step | ||
| 105 | + # therefore we do not do return here | ||
| 106 | + if dense_indice == -1: | ||
| 107 | + dense_boundary_col = -1 | ||
| 108 | + else: | ||
| 109 | + output_dense_ptr += dense_indice | ||
| 110 | + else: | ||
| 111 | + output_dense_ptr += pid * dense_matrix_stride | ||
| 112 | + | ||
| 113 | + if operation_function is not None: | ||
| 114 | + operation_jagged_value_ptr += begin * jagged_value_row_stride | ||
| 115 | + | ||
| 116 | + offset_row = tl.arange(0, thread_block_row_size) | ||
| 117 | + | ||
| 118 | + for _i in range(begin, end, thread_block_row_size): | ||
| 119 | + offset_col = tl.arange(0, thread_block_col_size) | ||
| 120 | + block_offset = offset_row[:, None] * dense_row_stride + offset_col[None, :] * dense_col_stride | ||
| 121 | + | ||
| 122 | + for _j in range(0, N, thread_block_col_size): | ||
| 123 | + dense_mask = (offset_row[:, None] < dense_boundary_row) & (offset_col[None, :] < dense_boundary_col) | ||
| 124 | + jagged_mask = (offset_row[:, None] < M) & (offset_col[None, :] < N) | ||
| 125 | + dense_values = tl.load(output_dense_ptr + block_offset, mask=dense_mask, other=0) | ||
| 126 | + if operation_function is not None: | ||
| 127 | + operation_jagged_value = tl.load(operation_jagged_value_ptr + block_offset, mask=jagged_mask, other=0) | ||
| 128 | + if operation_function == "add": | ||
| 129 | + dense_values = tensor_elementwise_add(dense_values, operation_jagged_value) | ||
| 130 | + else: | ||
| 131 | + dense_values = tensor_elementwise_mul(dense_values, operation_jagged_value) | ||
| 132 | + tl.store(jagged_value_ptr + block_offset, dense_values, mask=jagged_mask) | ||
| 133 | + offset_col += thread_block_col_size | ||
| 134 | + block_offset += thread_block_col_size | ||
| 135 | + offset_row += thread_block_row_size | ||
| 136 | + | ||
| 137 | + | ||
| 138 | +# this function parse the jagged tensor offsets to corresponding dense index position | ||
| 139 | +# to see the detail of it see the quip note : https://fb.quip.com/gnzpA7d13vqO | ||
| 140 | +# the FBGEMM implementation refer : https://www.internalfb.com/code/fbsource/[308212b2902c3182edcb5b204768321e032e8175]/fbcode/deeplearning/fbgemm/fbgemm_gpu/src/jagged_tensor_ops.cu?lines=280 | ||
| 141 | +# In FBGEMM it was computed by GPU but in triton currently has some compilation issue so we use CUP computation method as workaround | ||
| 142 | +# However in real-world case if we only dealing with 2d jagged tensor we don't need to use this function at all | ||
| 143 | +def _jagged_offsets_to_dense_indice( | ||
| 144 | + offsets: list[torch.Tensor], dense_strides: list[int], dense_sizes: list[int] | ||
| 145 | +) -> torch.Tensor: | ||
| 146 | + output_offset = torch.zeros(len(offsets[-1]) - 1, device="cpu", dtype=torch.int32) | ||
| 147 | + | ||
| 148 | + offsets_cpu = [] | ||
| 149 | + | ||
| 150 | + for offset in offsets: | ||
| 151 | + offsets_cpu.append(offset.cpu()) | ||
| 152 | + | ||
| 153 | + for i in range(0, len(offsets_cpu[-1]) - 1): | ||
| 154 | + idx = i | ||
| 155 | + result = 0 | ||
| 156 | + | ||
| 157 | + # flag to check if current offset is in the range of dense | ||
| 158 | + in_range = True | ||
| 159 | + for j in range(len(offsets_cpu) - 2, -1, -1): | ||
| 160 | + left = 0 | ||
| 161 | + right = offsets_cpu[j].size(0) | ||
| 162 | + | ||
| 163 | + # binary search found the corresponding offset group of current index | ||
| 164 | + while left < right: | ||
| 165 | + mid = left + (right - left) // 2 | ||
| 166 | + | ||
| 167 | + if offsets_cpu[j][mid] > idx: | ||
| 168 | + right = mid | ||
| 169 | + else: | ||
| 170 | + left = mid + 1 | ||
| 171 | + | ||
| 172 | + cur_val = idx - offsets_cpu[j][left - 1] | ||
| 173 | + | ||
| 174 | + if dense_sizes and cur_val >= dense_sizes[j + 1]: | ||
| 175 | + in_range = False | ||
| 176 | + break | ||
| 177 | + | ||
| 178 | + result += cur_val * dense_strides[j + 1] | ||
| 179 | + idx = left - 1 | ||
| 180 | + | ||
| 181 | + if in_range: | ||
| 182 | + result += idx * dense_strides[0] | ||
| 183 | + | ||
| 184 | + # another out of output dense range case | ||
| 185 | + if dense_sizes and idx > dense_sizes[0]: | ||
| 186 | + result = -1 | ||
| 187 | + output_offset[i] = result | ||
| 188 | + else: | ||
| 189 | + output_offset[i] = -1 | ||
| 190 | + | ||
| 191 | + return output_offset | ||
| @@ -0,0 +1,105 @@ | |||
| 1 | +# Copyright (c) Huawei Technologies Co., Ltd. 2026. All rights reserved. | ||
| 2 | +# | ||
| 3 | +# Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 4 | +# you may not use this file except in compliance with the License. | ||
| 5 | +# You may obtain a copy of the License at | ||
| 6 | +# | ||
| 7 | +# http://www.apache.org/licenses/LICENSE-2.0 | ||
| 8 | +# | ||
| 9 | +# Unless required by applicable law or agreed to in writing, software | ||
| 10 | +# distributed under the License is distributed on an "AS IS" BASIS, | ||
| 11 | +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| 12 | +# See the License for the specific language governing permissions and | ||
| 13 | +# limitations under the License. | ||
| 14 | +# ============================================================================== | ||
| 15 | +name: triton_dense_to_jagged | ||
| 16 | +api: pytorch | ||
| 17 | +version: v2.1 | ||
| 18 | +api_type: triton_dense_to_jagged | ||
| 19 | +triton_api_type: triton_dense_to_jagged | ||
| 20 | +triton_name: triton_dense_to_jagged.triton_dense_to_jagged | ||
| 21 | +generate: generate_triton_dense_to_jagged | ||
| 22 | +dtype_numbers: 3334 | ||
| 23 | +outputs: 0 | ||
| 24 | +standard: | ||
| 25 | + acc: single_bm | ||
| 26 | + perf: not_key | ||
| 27 | +inputs: | ||
| 28 | + - name: dense | ||
| 29 | + type: tensor | ||
| 30 | + required: true | ||
| 31 | + dtypes: | ||
| 32 | + values: [ fp32, fp16, bf16 ] | ||
| 33 | + ranges: | ||
| 34 | + valid: | ||
| 35 | + values: [ [-10000, 10000] ] | ||
| 36 | + random_types: | ||
| 37 | + - name: nd | ||
| 38 | + mean: [0, 5000] | ||
| 39 | + std: [1, 500] | ||
| 40 | + invalid: | ||
| 41 | + values: [ [-10000, 10000] ] | ||
| 42 | + random_types: | ||
| 43 | + - name: nd | ||
| 44 | + mean: [0, 5000] | ||
| 45 | + std: [1, 500] | ||
| 46 | + shapes: | ||
| 47 | + dim_numbers: | ||
| 48 | + values: [ 3, 4 ] | ||
| 49 | + dim_values: | ||
| 50 | + values: [ [8, 16], [16, 128], [128, 256] ] | ||
| 51 | + weights: [ 0.2, 0.2, 0.6 ] | ||
| 52 | + max_length: 4294967296 | ||
| 53 | + boundary: | ||
| 54 | + has_empty: false | ||
| 55 | + has_infnan: true | ||
| 56 | + has_scalar: false | ||
| 57 | + has_upper_border: false | ||
| 58 | + has_lower_border: false | ||
| 59 | + | ||
| 60 | + - name: jagged_offsets | ||
| 61 | + type: tensors | ||
| 62 | + required: true | ||
| 63 | + tuple_numbers: | ||
| 64 | + values: [ 1, 2 ] | ||
| 65 | + dtypes: | ||
| 66 | + values: [ int32, int64 ] | ||
| 67 | + shapes: | ||
| 68 | + dim_numbers: | ||
| 69 | + values: [ 1 ] | ||
| 70 | + dim_values: | ||
| 71 | + values: [ [2, 64] ] | ||
| 72 | + ranges: | ||
| 73 | + valid: | ||
| 74 | + values: [ [1, 128] ] | ||
| 75 | + invalid: | ||
| 76 | + values: [ [1, 128] ] | ||
| 77 | + boundary: | ||
| 78 | + has_empty: false | ||
| 79 | + has_infnan: false | ||
| 80 | + has_scalar: false | ||
| 81 | + has_upper_border: false | ||
| 82 | + has_lower_border: false | ||
| 83 | + | ||
| 84 | + - name: operation_function | ||
| 85 | + type: attr | ||
| 86 | + required: false | ||
| 87 | + dtypes: | ||
| 88 | + values: [ string ] | ||
| 89 | + ranges: | ||
| 90 | + valid: | ||
| 91 | + values: [ null, "add", "mul" ] | ||
| 92 | + weights: [ 0.2, 0.4, 0.4 ] | ||
| 93 | + invalid: | ||
| 94 | + values: [ null ] | ||
| 95 | + | ||
| 96 | + - name: operation_jagged_value | ||
| 97 | + type: scalar | ||
| 98 | + required: false | ||
| 99 | + dtypes: | ||
| 100 | + values: [ fp32, fp16, bf16 ] | ||
| 101 | + ranges: | ||
| 102 | + valid: | ||
| 103 | + values: [ [-5000, 5000] ] | ||
| 104 | + invalid: | ||
| 105 | + values: [ [-5000, 5000] ] | ||
Aexperimental/triton/atk_test/triton_dense_to_jagged/npu_atk_test/triton_dense_to_jagged_api.py+343-0
| @@ -0,0 +1,343 @@ | |||
| 1 | +# Copyright (c) Huawei Technologies Co., Ltd. 2026. All rights reserved. | ||
| 2 | +# | ||
| 3 | +# Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 4 | +# you may not use this file except in compliance with the License. | ||
| 5 | +# You may obtain a copy of the License at | ||
| 6 | +# | ||
| 7 | +# http://www.apache.org/licenses/LICENSE-2.0 | ||
| 8 | +# | ||
| 9 | +# Unless required by applicable law or agreed to in writing, software | ||
| 10 | +# distributed under the License is distributed on an "AS IS" BASIS, | ||
| 11 | +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| 12 | +# See the License for the specific language governing permissions and | ||
| 13 | +# limitations under the License. | ||
| 14 | +# ============================================================================== | ||
| 15 | + | ||
| 16 | +import torch | ||
| 17 | + | ||
| 18 | +from atk.configs.dataset_config import InputDataset | ||
| 19 | +from atk.tasks.api_execute import register | ||
| 20 | +from atk.tasks.api_execute.base_api import BaseApi | ||
| 21 | + | ||
| 22 | +# 从被测代码直接引用 triton kernel | ||
| 23 | +from triton_dense_to_jagged import triton_dense_to_jagged, _jagged_offsets_to_dense_indice | ||
| 24 | + | ||
| 25 | + | ||
| 26 | +def get_device(): | ||
| 27 | + """获取可用的设备,优先 NPU,其次 CUDA,最后 CPU""" | ||
| 28 | + if hasattr(torch, 'npu') and torch.npu.is_available(): | ||
| 29 | + return 'npu' | ||
| 30 | + try: | ||
| 31 | + if torch.cuda.is_available(): | ||
| 32 | + return 'cuda' | ||
| 33 | + except (RuntimeError, AttributeError): | ||
| 34 | + # CUDA driver/runtime 未安装或属性不可用时回退到 CPU | ||
| 35 | + pass | ||
| 36 | + return 'cpu' | ||
| 37 | + | ||
| 38 | + | ||
| 39 | +def _generate_jagged_offsets(dense: torch.Tensor, jagged_offsets_input: list, device: str) -> list: | ||
| 40 | + """根据 dense 的 shape 生成满足约束的 jagged_offsets | ||
| 41 | + | ||
| 42 | + 这个函数用于在 API 初始化时,根据 dense 的实际 shape 生成符合约束的 offset。 | ||
| 43 | + 使用确定性方式生成 offset,确保不同机器上结果一致。 | ||
| 44 | + | ||
Z | |||
| 45 | + 参数: | ||
| 46 | + dense: 输入的 dense tensor | ||
| 47 | + jagged_offsets_input: 从 yaml 生成的原始 offset 列表 | ||
| 48 | + device: 目标设备('npu'/'cuda'/'cpu') | ||
| 49 | + | ||
| 50 | + 返回: | ||
| 51 | + 调整后的 offset 列表 | ||
| 52 | + """ | ||
| 53 | + dense_shape = list(dense.shape) | ||
| 54 | + dense_dim = len(dense_shape) | ||
| 55 | + | ||
| 56 | + adjusted_offsets = [] | ||
| 57 | + | ||
| 58 | + if dense_dim == 3: | ||
| 59 | + # ==================== 3D dense [B, seq, D] ==================== | ||
| 60 | + # 对应 1 个 offset,shape [B+1] | ||
| 61 | + # 约束:2*B <= offset[-1] <= B*seq,相邻差值不超过 seq | ||
| 62 | + B = dense_shape[0] | ||
| 63 | + seq = dense_shape[1] | ||
| 64 | + | ||
| 65 | + # 参考 offset 仅用于取 dtype/device;其数值完全由下面按 dense 维度重建, | ||
| 66 | + # 因此不依赖输入的 offset 层数。3D dense 必须且只会生成 1 层 offset。 | ||
| 67 | + if jagged_offsets_input: | ||
| 68 | + ref_offset = jagged_offsets_input[0].to(device) | ||
| 69 | + else: | ||
| 70 | + ref_offset = torch.zeros(1, dtype=torch.int64, device=device) | ||
| 71 | + | ||
| 72 | + # 计算目标最后一个值:取中间值 | ||
| 73 | + min_last = 2 * B | ||
| 74 | + max_last = B * seq | ||
| 75 | + target_last = (min_last + max_last) // 2 | ||
| 76 | + | ||
| 77 | + # 每个元素的平均长度 | ||
| 78 | + avg_len = max(target_last // B, 1) | ||
| 79 | + avg_len = min(avg_len, seq) # 不超过 seq | ||
| 80 | + | ||
| 81 | + # 分配长度 | ||
| 82 | + remainder = target_last - avg_len * B | ||
| 83 | + segment_lengths = [] | ||
| 84 | + for i in range(B): | ||
| 85 | + if i < remainder: | ||
| 86 | + seg_len = avg_len + 1 | ||
| 87 | + else: | ||
| 88 | + seg_len = avg_len | ||
| 89 | + segment_lengths.append(seg_len) | ||
| 90 | + | ||
| 91 | + cumulative = [0] | ||
| 92 | + for seg_len in segment_lengths: | ||
| 93 | + cumulative.append(cumulative[-1] + seg_len) | ||
| 94 | + | ||
| 95 | + offset0 = torch.tensor(cumulative, dtype=ref_offset.dtype, device=device) | ||
| 96 | + adjusted_offsets.append(offset0) | ||
| 97 | + | ||
| 98 | + elif dense_dim == 4: | ||
| 99 | + # ==================== 4D dense [B1, B2, seq, D] ==================== | ||
| 100 | + # 对应 2 个 offset | ||
| 101 | + # offset0: shape [B1+1] | ||
| 102 | + # offset1: shape [B1*B2+1] | ||
| 103 | + B1 = dense_shape[0] | ||
| 104 | + B2 = dense_shape[1] | ||
| 105 | + seq = dense_shape[2] | ||
| 106 | + | ||
| 107 | + # 参考 offset 仅用于取 dtype/device;其数值完全由下面按 dense 维度重建, | ||
| 108 | + # 因此不依赖输入的 offset 层数。4D dense 必须且只会生成 2 层 offset, | ||
| 109 | + # 保证 len(adjusted_offsets) == dense.ndim - 2 == 2。 | ||
| 110 | + if jagged_offsets_input: | ||
| 111 | + ref_offset = jagged_offsets_input[0].to(device) | ||
| 112 | + else: | ||
| 113 | + # 兜底:jagged_offsets 输入为空时仍能按 int64 构建 2 层 offset | ||
| 114 | + ref_offset = torch.zeros(1, dtype=torch.int64, device=device) | ||
| 115 | + | ||
| 116 | + # 第一个 offset:对应 B1 维度 | ||
| 117 | + # 约束:jagged_offsets[0][-1] = B1*B2,相邻差值不超过 B2 | ||
| 118 | + # 每个 batch 元素应该有 B2 个元素 | ||
| 119 | + segment_lengths = [B2] * B1 | ||
| 120 | + cumulative = [0] | ||
| 121 | + for seg_len in segment_lengths: | ||
| 122 | + cumulative.append(cumulative[-1] + seg_len) | ||
| 123 | + | ||
| 124 | + offset0 = torch.tensor(cumulative, dtype=ref_offset.dtype, device=device) | ||
| 125 | + adjusted_offsets.append(offset0) | ||
| 126 | + | ||
| 127 | + # 第二个 offset:对应 B1*B2 维度 | ||
| 128 | + # 约束:2*B1*B2 <= jagged_offsets[1][-1] <= B1*B2*seq,相邻差值不超过 seq | ||
| 129 | + total_inner = B1 * B2 | ||
| 130 | + # 取中间值:2*B1*B2 和 B1*B2*seq 的中间 | ||
| 131 | + min_last = 2 * total_inner | ||
| 132 | + max_last = total_inner * seq | ||
| 133 | + target_last = (min_last + max_last) // 2 | ||
| 134 | + | ||
| 135 | + # 每个元素的平均长度 | ||
| 136 | + avg_len = max(target_last // total_inner, 1) | ||
| 137 | + avg_len = min(avg_len, seq) # 不超过 seq | ||
| 138 | + | ||
| 139 | + # 分配长度:前几个元素多分配1 | ||
| 140 | + remainder = target_last - avg_len * total_inner | ||
| 141 | + segment_lengths = [] | ||
| 142 | + for i in range(total_inner): | ||
| 143 | + if i < remainder: | ||
| 144 | + seg_len = avg_len + 1 | ||
| 145 | + else: | ||
| 146 | + seg_len = avg_len | ||
| 147 | + segment_lengths.append(seg_len) | ||
| 148 | + | ||
| 149 | + cumulative = [0] | ||
| 150 | + for seg_len in segment_lengths: | ||
| 151 | + cumulative.append(cumulative[-1] + seg_len) | ||
| 152 | + | ||
| 153 | + offset1 = torch.tensor(cumulative, dtype=ref_offset.dtype, device=device) | ||
| 154 | + adjusted_offsets.append(offset1) | ||
| 155 | + else: | ||
| 156 | + # 其他维度,直接使用 | ||
| 157 | + adjusted_offsets = [offset.to(device) for offset in jagged_offsets_input] | ||
| 158 | + | ||
| 159 | + # 确保第一个元素为 0(偏移从 0 开始) | ||
| 160 | + for offset in adjusted_offsets: | ||
| 161 | + if offset[0].item() != 0: | ||
| 162 | + offset[0] = torch.tensor(0, dtype=offset.dtype, device=device) | ||
| 163 | + | ||
| 164 | + return adjusted_offsets | ||
| 165 | + | ||
| 166 | + | ||
| 167 | + | ||
| 168 | +class TritonDenseToJaggedApi(BaseApi): | ||
| 169 | + """dense_to_jagged 算子的 ATK API 封装 | ||
| 170 | + | ||
| 171 | + 继承 BaseApi,实现算子的初始化和执行。 | ||
| 172 | + """ | ||
| 173 | + | ||
| 174 | + # 在类层级预声明属性,便于 pylint 静态检查; | ||
| 175 | + # 实际值在 init_by_input_data 中由 ATK 框架注入。 | ||
| 176 | + _device = None | ||
| 177 | + _output_jagged_value = None | ||
| 178 | + _jagged_offsets = None | ||
| 179 | + _grid = None | ||
| 180 | + _JAGGED_DIM = None | ||
| 181 | + _dense_indices = None | ||
| 182 | + _dense = None | ||
| 183 | + _dense_col_stride = None | ||
| 184 | + _dense_row_stride = None | ||
| 185 | + _dense_matrix_stride = None | ||
| 186 | + _operation_function = None | ||
| 187 | + _operation_jagged_values = None | ||
| 188 | + | ||
| 189 | + def init_by_input_data(self, input_data: InputDataset): | ||
| 190 | + """ | ||
| 191 | + 初始化操作 | ||
| 192 | + | ||
| 193 | + 在调用算子前执行一次,用于: | ||
| 194 | + 1. 获取输入参数 | ||
| 195 | + 2. 生成满足约束的 jagged_offsets | ||
| 196 | + 3. 准备输出 tensor | ||
| 197 | + 4. 计算 kernel 所需的参数 | ||
| 198 | + """ | ||
| 199 | + device = get_device() | ||
| 200 | + | ||
| 201 | + # 同步操作,确保之前的所有操作完成 | ||
| 202 | + if device == 'npu': | ||
| 203 | + torch.npu.synchronize() | ||
| 204 | + elif device == 'cuda': | ||
| 205 | + torch.cuda.synchronize() | ||
| 206 | + | ||
| 207 | + # ==================== 获取输入参数 ==================== | ||
| 208 | + # 从 InputDataset 中获取 kwargs 形式的参数 | ||
| 209 | + dense_input = input_data.kwargs.get("dense", None) | ||
| 210 | + jagged_offsets_input = input_data.kwargs.get("jagged_offsets", []) | ||
| 211 | + operation_function_input = input_data.kwargs.get("operation_function", None) | ||
| 212 | + operation_jagged_value_input = input_data.kwargs.get("operation_jagged_value", None) | ||
| 213 | + | ||
| 214 | + if dense_input is None: | ||
| 215 | + raise ValueError("dense not found") | ||
| 216 | + | ||
| 217 | + # ==================== 转换 dense 到正确设备 ==================== | ||
| 218 | + if not isinstance(dense_input, torch.Tensor): | ||
| 219 | + dense = torch.tensor(dense_input, dtype=torch.float32, device=device) | ||
| 220 | + else: | ||
| 221 | + dense = dense_input.to(device) | ||
| 222 | + | ||
| 223 | + # ==================== 处理 jagged_offsets ==================== | ||
| 224 | + jagged_offsets = [] | ||
| 225 | + for offset_tensor in jagged_offsets_input: | ||
| 226 | + if isinstance(offset_tensor, torch.Tensor): | ||
| 227 | + jagged_offsets.append(offset_tensor.to(device)) | ||
| 228 | + else: | ||
| 229 | + jagged_offsets.append(torch.tensor(offset_tensor, device=device)) | ||
| 230 | + | ||
| 231 | + # ==================== 生成满足约束的 jagged_offsets ==================== | ||
| 232 | + jagged_offsets = _generate_jagged_offsets(dense, jagged_offsets, device) | ||
| 233 | + | ||
| 234 | + # ==================== 计算输出的 jagged tensor 的大小 ==================== | ||
| 235 | + # 最后一个 offset 的最后一个值就是输出的行数 | ||
| 236 | + actual_total_valid = jagged_offsets[-1][-1].item() | ||
| 237 | + | ||
| 238 | + # ==================== 处理 operation_function ==================== | ||
| 239 | + # operation_function 可以是 null, "add", 或 "mul" | ||
| 240 | + operation_function = None | ||
| 241 | + op_func_str = str(operation_function_input) if operation_function_input is not None else "" | ||
| 242 | + if op_func_str.lower() not in ['null', 'none', '']: | ||
| 243 | + operation_function = op_func_str | ||
| 244 | + | ||
| 245 | + # ==================== 处理 operation_jagged_values ==================== | ||
| 246 | + # 当 operation_function 不为 null 时,需要生成 operation_jagged_values | ||
| 247 | + # 在确认输出的 jagged_tensor 的 shape 之后,再生成相同 shape, | ||
| 248 | + # 值都为 operation_jagged_value 的 tensor | ||
| 249 | + operation_jagged_values = None | ||
| 250 | + if operation_function is not None and operation_function not in ['null', 'None', '']: | ||
| 251 | + if operation_jagged_value_input is None: | ||
| 252 | + # 关闭融合以避免 kernel 对 None 指针做算术 | ||
| 253 | + operation_function = None | ||
| 254 | + else: | ||
| 255 | + op_value = float(operation_jagged_value_input) | ||
| 256 | + # 生成与 output_jagged_value 相同 shape 的 tensor,值都为 op_value | ||
| 257 | + # 数据类型与 dense 保持一致 | ||
| 258 | + operation_jagged_values = torch.full( | ||
| 259 | + (int(actual_total_valid), dense.size(-1)), op_value, dtype=dense.dtype, device=device | ||
| 260 | + ) | ||
| 261 | + | ||
| 262 | + # ==================== 创建输出 tensor ==================== | ||
| 263 | + output_jagged_value = torch.empty( | ||
| 264 | + (int(actual_total_valid), dense.size(-1)), | ||
| 265 | + device=device, | ||
| 266 | + dtype=dense.dtype, | ||
| 267 | + ) | ||
| 268 | + | ||
| 269 | + # ==================== 计算 kernel 参数 ==================== | ||
| 270 | + # grid: triton kernel 的启动配置,等于 offset 的元素数减 1 | ||
| 271 | + grid = (jagged_offsets[-1].size(0) - 1,) | ||
| 272 | + | ||
| 273 | + # JAGGED_DIM: jagged tensor 的维度,等于 offset 数量 + 1 | ||
| 274 | + JAGGED_DIM = len(jagged_offsets) + 1 | ||
| 275 | + | ||
| 276 | + dense_indices = None | ||
| 277 | + if len(jagged_offsets) > 1: | ||
| 278 | + dense_indices = _jagged_offsets_to_dense_indice( | ||
| 279 | + jagged_offsets, | ||
| 280 | + dense.stride()[:-2], | ||
| 281 | + dense.size()[:-2], | ||
| 282 | + ) | ||
| 283 | + if device == 'npu': | ||
| 284 | + dense_indices = dense_indices.npu() | ||
| 285 | + elif device == 'cuda': | ||
| 286 | + dense_indices = dense_indices.cuda() | ||
| 287 | + | ||
| 288 | + # 获取 dense 的 stride 信息,用于 kernel 计算 | ||
| 289 | + dense_col_stride = dense.stride(-1) # 最后一维的 stride | ||
| 290 | + dense_row_stride = dense.stride(-2) # 倒数第二维的 stride | ||
| 291 | + dense_matrix_stride = dense.stride(-3) # 倒数第三维的 stride | ||
| 292 | + | ||
| 293 | + # ==================== 保存实例变量 ==================== | ||
| 294 | + # 保存所有需要用到变量,供 __call__ 使用 | ||
| 295 | + self._device = device | ||
| 296 | + self._output_jagged_value = output_jagged_value | ||
| 297 | + self._jagged_offsets = jagged_offsets | ||
| 298 | + self._grid = grid | ||
| 299 | + self._JAGGED_DIM = JAGGED_DIM | ||
| 300 | + self._dense_indices = dense_indices | ||
| 301 | + self._dense = dense | ||
| 302 | + self._dense_col_stride = dense_col_stride | ||
| 303 | + self._dense_row_stride = dense_row_stride | ||
| 304 | + self._dense_matrix_stride = dense_matrix_stride | ||
| 305 | + self._operation_function = operation_function | ||
| 306 | + self._operation_jagged_values = operation_jagged_values | ||
| 307 | + | ||
| 308 | + def __call__(self, input_data: InputDataset, with_output: bool = False): | ||
| 309 | + """ | ||
| 310 | + 执行 dense_to_jagged 算子 | ||
| 311 | + | ||
| 312 | + 每次测试用例执行时都会调用这个函数。 | ||
| 313 | + """ | ||
| 314 | + device = self._device | ||
| 315 | + | ||
| 316 | + # 同步操作 | ||
| 317 | + if device == 'npu': | ||
| 318 | + torch.npu.synchronize() | ||
| 319 | + elif device == 'cuda': | ||
| 320 | + torch.cuda.synchronize() | ||
| 321 | + | ||
| 322 | + # 调用 triton kernel | ||
| 323 | + triton_dense_to_jagged[self._grid]( | ||
| 324 | + self._output_jagged_value, # 输出:jagged tensor | ||
| 325 | + self._jagged_offsets[-1], # 输入:offsets | ||
| 326 | + self._output_jagged_value.stride(0), # 输出 tensor 的 stride | ||
| 327 | + self._dense, # 输入:dense tensor | ||
| 328 | + self._dense_indices, # dense indices | ||
| 329 | + self._dense_col_stride, # dense 最后一维 stride | ||
| 330 | + self._dense_row_stride, # dense 倒数第二维 stride | ||
| 331 | + self._dense_matrix_stride, # dense 倒数第三维 stride | ||
| 332 | + self._JAGGED_DIM, # jagged 维度数 | ||
| 333 | + operation_function=self._operation_function, | ||
| 334 | + operation_jagged_value_ptr=self._operation_jagged_values, | ||
| 335 | + ) | ||
| 336 | + | ||
| 337 | + # 同步操作 | ||
| 338 | + if device == 'npu': | ||
| 339 | + torch.npu.synchronize() | ||
| 340 | + elif device == 'cuda': | ||
| 341 | + torch.cuda.synchronize() | ||
| 342 | + | ||
| 343 | + return self._output_jagged_value | ||
[P2] [P2] 被测入参被计算覆盖(缺陷模式 1.19):
_generate_jagged_offsets(L67)仅使用传入 offsets 的 shape(len),完全按内部规则重新生成 offset 数值,yaml 配置的 jagged_offsets 实际取值被丢弃。后果:① 任何针对非法/边界 offset 的 invalid 用例实际不生效;② 截断、padding、dense_indice==-1 跳过等分支无法被真实触发。若确为有意改为自动生成,请在 PR 描述与 yaml 注释中明确,并补充“直接传入 offsets”的用例以覆盖上述分支。