已合并
[feat] 新增triton_dense_to_jagged算子atk测试迁移(npu部分) #123
[feat] 新增triton_dense_to_jagged算子atk测试迁移(npu部分) #123
已合并
QianZH97创建于 27 天前
共 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
@@ -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+@GENERATOR_REGISTRY.register("generate_triton_dense_to_jagged")
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+@triton.jit
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+@triton.jit
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+@triton.autotune(
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+@triton.jit
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] ]
@@ -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
Zzengxiong27 天前

[P2] [P2] 被测入参被计算覆盖(缺陷模式 1.19):_generate_jagged_offsets(L67)仅使用传入 offsets 的 shape(len),完全按内部规则重新生成 offset 数值,yaml 配置的 jagged_offsets 实际取值被丢弃。后果:① 任何针对非法/边界 offset 的 invalid 用例实际不生效;② 截断、padding、dense_indice==-1 跳过等分支无法被真实触发。若确为有意改为自动生成,请在 PR 描述与 yaml 注释中明确,并补充“直接传入 offsets”的用例以覆盖上述分支。

likedislike
QianZH97
QianZH97
27 天前 评论:
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+@register("triton_dense_to_jagged")
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