Triton算子精度评估技能
1. 功能概述
该技能用于自动化评估Triton算子实现的精度,通过与PyTorch(CPU或NPU)的对应算子实现进行比对,生成详细的精度验证报告。
核心功能
- 自动接收Triton算子实现
- 支持与CPU或NPU上的Torch小算子进行比对
- 支持多种数据类型(float16、float32、int8、uint8等)
- 自动生成精度验证报告
- 支持批量测试不同参数配置
2. 调用时机
在以下情况下调用此技能:
- 需要验证Triton算子实现的精度正确性
- 需要与PyTorch原生算子进行精度比对
- 需要生成标准化的精度验证报告
- 需要批量测试Triton算子在不同数据类型和参数下的精度表现
3. 实现原理
3.1 工作流程
┌─────────────────┐ ┌─────────────────┐ ┌─────────────────┐
│ Triton算子实现 │────▶│ 生成测试数据 │────▶│ 执行Torch对比实现 │
└─────────────────┘ └─────────────────┘ └─────────────────┘
▲ │ │
│ ▼ ▼
│ ┌─────────────────┐ ┌─────────────────┐
│ │ 执行Triton实现 │ │ 计算误差指标 │
│ └─────────────────┘ └─────────────────┘
│ │ │
└─────────────────────┼─────────────────────┘
│
▼
┌─────────────────┐
│ 生成精度报告 │
└─────────────────┘
3.2 核心组件
- 测试数据生成:使用
test_common.generate_numpy() 生成随机测试数据
- Torch对比实现:用户提供的Torch算子实现(如示例中的
torch_pointwise())
- Triton算子执行:使用Triton JIT编译并执行用户提供的Triton kernel
- 精度验证:使用
test_common.validate_cmp() 进行精度比对,支持不同数据类型的误差阈值
- 报告生成:生成包含误差指标的精度验证报告
4. 使用方法
4.1 前置条件
- 已安装Triton和PyTorch环境
- 已安装
torch_npu(如果使用NPU进行测试)
- 已准备Triton算子实现代码
4.2 编写测试用例
创建测试文件(如 test_abs.py),包含以下内容:
导入必要模块:
import triton
import triton.language as tl
import numpy as np
import torch
import pytest
import test_common
实现Torch对比函数:
def torch_pointwise(x0):
# 实现与Triton算子对应的Torch功能
return torch.abs(x0)
实现Triton算子:
@triton.jit
def triton_abs(in_ptr0, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr):
# Triton kernel实现
offset = tl.program_id(0) * XBLOCK
base1 = tl.arange(0, XBLOCK_SUB)
loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB
for loop1 in range(loops1):
x0_prime = offset + (loop1 * XBLOCK_SUB) + base1
x0 = offset + (loop1 * XBLOCK_SUB) + base1
tmp0 = tl.load(in_ptr0 + (x0), None)
tmp2 = tl.abs(tmp0)
tl.store(out_ptr0 + (x0), tmp2, None)
编写测试用例:
@pytest.mark.parametrize('param_list',
[
['float16', (2, 4096, 8), 32, 2048, 64],
['float32', (2, 4096, 8), 32, 2048, 64],
['int8', (2, 4096, 8), 32, 2048, 64],
['uint8', (2, 4096, 8), 32, 2048, 64],
]
)
def test_case(param_list):
dtype, shape, ncore, xblock, xblock_sub = param_list
np_x0 = test_common.generate_numpy(shape, dtype)
x0 = torch.from_numpy(np_x0).to(eval('torch.' + dtype)).npu()
y_ref = torch_pointwise(x0)
y_cal = torch.zeros(shape, dtype = eval('torch.' + dtype)).npu()
triton_abs[ncore, 1, 1](x0, y_cal, xblock, xblock_sub)
test_common.validate_cmp(dtype, y_cal, y_ref)
4.3 运行测试
# 运行单个测试文件
pytest test_abs.py -v
# 运行所有测试文件
pytest ./examples/ -v
5. 精度验证规则
5.1 不同数据类型的验证规则
| 数据类型 |
验证方式 |
误差阈值 |
| float16 |
相对误差 |
rtol=1e-03, atol=1e-03 |
| float32 |
相对误差 |
rtol=1e-04, atol=1e-04 |
| bfloat16 |
相对误差 |
rtol=1e-03, atol=1e-03 |
| int32/int64/int16/int8 |
完全相等 |
- |
| uint32/uint64/uint16/uint8 |
完全相等 |
- |
| bool |
完全相等 |
- |
5.2 误差指标
- 平均相对误差(MERE):所有元素相对误差的平均值
- 最大相对误差(MARE):所有元素相对误差的最大值
- 绝对误差:元素值之差的绝对值
6. 精度报告格式
生成的精度报告(如 eco_report.txt)包含以下内容:
================================================================================
Triton算子精度验证报告
--------------------------------------------------------------------------------
[验证配置]:
数据类型: float32 (Single Precision)
MERE阈值: 1.220703e-04
MARE阈值: 1.220703e-03 (10×MERE阈值)
小值域阈值: 1.000000e-07
--------------------------------------------------------------------------------
[精度标准]:
float16: 相对误差 rtol=1e-03, atol=1e-03
float32: 相对误差 rtol=1e-04, atol=1e-04
bfloat16: 相对误差 rtol=1e-03, atol=1e-03
int32/int64/int16/int8: 完全相等
uint32/uint64/uint16/uint8: 完全相等
bool: 完全相等
--------------------------------------------------------------------------------
[验证结果]:
验证结果: FAIL
样本总数: 4096
--------------------------------------------------------------------------------
[误差指标]:
平均相对误差(MERE): 6.642197e-03
阈值要求: MERE < 1.220703e-04
最大相对误差(MARE): 3.458786e+00
阈值要求: MARE < 1.220703e-03
--------------------------------------------------------------------------------
[判定条件]:
✓ MERE < 阈值: False
✓ MARE < 10×阈值: False
✓ 总体结果: False
================================================================================
6.1 报告内容要求
精度报告必须包含以下内容:
- 验证配置:算子名称、测试形状、数据类型、NPU核心数等
- 精度标准:每个数据类型的具体精度要求(误差阈值或完全相等)
- 验证结果:测试总数、通过数量、失败数量、总体结果
- 详细误差指标:每个数据类型的平均相对误差、最大相对误差、最大绝对误差
- 判定条件:所有数据类型测试通过的状态
其中,精度标准部分必须列出所有支持的数据类型及其对应的精度要求,确保报告的可读性和可追溯性。
7. 示例
7.1 输入示例(Triton算子实现)
@triton.jit
def triton_abs(in_ptr0, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr):
offset = tl.program_id(0) * XBLOCK
base1 = tl.arange(0, XBLOCK_SUB)
loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB
for loop1 in range(loops1):
x0_prime = offset + (loop1 * XBLOCK_SUB) + base1
x0 = offset + (loop1 * XBLOCK_SUB) + base1
tmp0 = tl.load(in_ptr0 + (x0), None)
tmp2 = tl.abs(tmp0)
tl.store(out_ptr0 + (x0), tmp2, None)
7.2 输出示例(精度报告)
================================================================================
Triton算子精度验证报告
--------------------------------------------------------------------------------
[验证配置]:
数据类型: float32 (Single Precision)
MERE阈值: 1.220703e-04
MARE阈值: 1.220703e-03 (10×MERE阈值)
小值域阈值: 1.000000e-07
--------------------------------------------------------------------------------
[验证结果]:
验证结果: PASS
样本总数: 4096
--------------------------------------------------------------------------------
[误差指标]:
平均相对误差(MERE): 0.000000e+00
阈值要求: MERE < 1.220703e-04
最大相对误差(MARE): 0.000000e+00
阈值要求: MARE < 1.220703e-03
--------------------------------------------------------------------------------
[判定条件]:
✓ MERE < 阈值: True
✓ MARE < 10×阈值: True
✓ 总体结果: True
================================================================================
8. 注意事项
- 环境配置:确保已正确安装Triton、PyTorch和必要的依赖
- 数据类型:不同数据类型有不同的精度验证规则,请根据算子类型选择合适的数据类型
- 参数配置:根据算子复杂度和硬件资源调整测试参数(如ncore、xblock等)
- 错误处理:如果测试失败,检查Triton算子实现和Torch对比实现是否功能一致
- 报告解读:根据精度报告中的误差指标调整算子实现以提高精度
9. 故障处理
| 问题 |
可能原因 |
解决方案 |
| Triton kernel编译失败 |
Triton语法错误或版本不兼容 |
检查Triton语法,确保Triton版本与代码兼容 |
| 精度验证失败 |
算子实现逻辑错误或精度损失 |
检查算子实现,调整算法以提高精度 |
| NPU设备不可用 |
未安装torch_npu或设备未正确配置 |
安装torch_npu,检查NPU设备状态 |
| 内存不足 |
测试数据过大 |
减小测试数据规模或调整参数配置 |
10. 参考资料
1---2name: triton-operator-precision-eval3description: 接收Triton算子实现,自动调用Torch小算子实现(CPU或NPU)进行精度比对,并生成精度报告。用于验证Triton算子实现的正确性和精度。4---56# Triton算子精度评估技能78---910## 1. 功能概述1112该技能用于自动化评估Triton算子实现的精度,通过与PyTorch(CPU或NPU)的对应算子实现进行比对,生成详细的精度验证报告。1314### 核心功能15- 自动接收Triton算子实现16- 支持与CPU或NPU上的Torch小算子进行比对17- 支持多种数据类型(float16、float32、int8、uint8等)18- 自动生成精度验证报告19- 支持批量测试不同参数配置2021---2223## 2. 调用时机2425在以下情况下调用此技能:26- 需要验证Triton算子实现的精度正确性27- 需要与PyTorch原生算子进行精度比对28- 需要生成标准化的精度验证报告29- 需要批量测试Triton算子在不同数据类型和参数下的精度表现3031---3233## 3. 实现原理3435### 3.1 工作流程3637```38┌─────────────────┐ ┌─────────────────┐ ┌─────────────────┐39│ Triton算子实现 │────▶│ 生成测试数据 │────▶│ 执行Torch对比实现 │40└─────────────────┘ └─────────────────┘ └─────────────────┘41 ▲ │ │42 │ ▼ ▼43 │ ┌─────────────────┐ ┌─────────────────┐44 │ │ 执行Triton实现 │ │ 计算误差指标 │45 │ └─────────────────┘ └─────────────────┘46 │ │ │47 └─────────────────────┼─────────────────────┘48 │49 ▼50 ┌─────────────────┐51 │ 生成精度报告 │52 └─────────────────┘53```5455### 3.2 核心组件56571. **测试数据生成**:使用 `test_common.generate_numpy()` 生成随机测试数据582. **Torch对比实现**:用户提供的Torch算子实现(如示例中的 `torch_pointwise()`)593. **Triton算子执行**:使用Triton JIT编译并执行用户提供的Triton kernel604. **精度验证**:使用 `test_common.validate_cmp()` 进行精度比对,支持不同数据类型的误差阈值615. **报告生成**:生成包含误差指标的精度验证报告6263---6465## 4. 使用方法6667### 4.1 前置条件6869- 已安装Triton和PyTorch环境70- 已安装 `torch_npu`(如果使用NPU进行测试)71- 已准备Triton算子实现代码7273### 4.2 编写测试用例7475创建测试文件(如 `test_abs.py`),包含以下内容:76771. **导入必要模块**:78 ```python79 import triton80 import triton.language as tl81 import numpy as np82 import torch83 import pytest84 import test_common85 ```86872. **实现Torch对比函数**:88 ```python89 def torch_pointwise(x0):90 # 实现与Triton算子对应的Torch功能91 return torch.abs(x0)92 ```93943. **实现Triton算子**:95 ```python96 @triton.jit97 def triton_abs(in_ptr0, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr):98 # Triton kernel实现99 offset = tl.program_id(0) * XBLOCK100 base1 = tl.arange(0, XBLOCK_SUB)101 loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB102 for loop1 in range(loops1):103 x0_prime = offset + (loop1 * XBLOCK_SUB) + base1104 x0 = offset + (loop1 * XBLOCK_SUB) + base1105 tmp0 = tl.load(in_ptr0 + (x0), None)106 tmp2 = tl.abs(tmp0)107 tl.store(out_ptr0 + (x0), tmp2, None)108 ```1091104. **编写测试用例**:111 ```python112 @pytest.mark.parametrize('param_list',113 [114 ['float16', (2, 4096, 8), 32, 2048, 64],115 ['float32', (2, 4096, 8), 32, 2048, 64],116 ['int8', (2, 4096, 8), 32, 2048, 64],117 ['uint8', (2, 4096, 8), 32, 2048, 64],118 ]119 )120121 def test_case(param_list):122 dtype, shape, ncore, xblock, xblock_sub = param_list123 np_x0 = test_common.generate_numpy(shape, dtype)124 x0 = torch.from_numpy(np_x0).to(eval('torch.' + dtype)).npu()125 y_ref = torch_pointwise(x0)126 y_cal = torch.zeros(shape, dtype = eval('torch.' + dtype)).npu()127 triton_abs[ncore, 1, 1](x0, y_cal, xblock, xblock_sub)128 test_common.validate_cmp(dtype, y_cal, y_ref)129 ```130131### 4.3 运行测试132133```bash134# 运行单个测试文件135pytest test_abs.py -v136137# 运行所有测试文件138pytest ./examples/ -v139```140141---142143## 5. 精度验证规则144145### 5.1 不同数据类型的验证规则146147| 数据类型 | 验证方式 | 误差阈值 |148|---------|---------|---------|149| float16 | 相对误差 | rtol=1e-03, atol=1e-03 |150| float32 | 相对误差 | rtol=1e-04, atol=1e-04 |151| bfloat16 | 相对误差 | rtol=1e-03, atol=1e-03 |152| int32/int64/int16/int8 | 完全相等 | - |153| uint32/uint64/uint16/uint8 | 完全相等 | - |154| bool | 完全相等 | - |155156### 5.2 误差指标157158- **平均相对误差(MERE)**:所有元素相对误差的平均值159- **最大相对误差(MARE)**:所有元素相对误差的最大值160- **绝对误差**:元素值之差的绝对值161162---163164## 6. 精度报告格式165166生成的精度报告(如 `eco_report.txt`)包含以下内容:167168```169================================================================================170 Triton算子精度验证报告 171--------------------------------------------------------------------------------172[验证配置]:173 数据类型: float32 (Single Precision)174 MERE阈值: 1.220703e-04175 MARE阈值: 1.220703e-03 (10×MERE阈值)176 小值域阈值: 1.000000e-07177--------------------------------------------------------------------------------178[精度标准]:179 float16: 相对误差 rtol=1e-03, atol=1e-03180 float32: 相对误差 rtol=1e-04, atol=1e-04181 bfloat16: 相对误差 rtol=1e-03, atol=1e-03182 int32/int64/int16/int8: 完全相等183 uint32/uint64/uint16/uint8: 完全相等184 bool: 完全相等185--------------------------------------------------------------------------------186[验证结果]:187 验证结果: FAIL188 样本总数: 4096189--------------------------------------------------------------------------------190[误差指标]:191 平均相对误差(MERE): 6.642197e-03192 阈值要求: MERE < 1.220703e-04193 最大相对误差(MARE): 3.458786e+00194 阈值要求: MARE < 1.220703e-03195--------------------------------------------------------------------------------196[判定条件]:197 ✓ MERE < 阈值: False198 ✓ MARE < 10×阈值: False199 ✓ 总体结果: False200================================================================================201```202203### 6.1 报告内容要求204205精度报告必须包含以下内容:2062071. **验证配置**:算子名称、测试形状、数据类型、NPU核心数等2082. **精度标准**:每个数据类型的具体精度要求(误差阈值或完全相等)2093. **验证结果**:测试总数、通过数量、失败数量、总体结果2104. **详细误差指标**:每个数据类型的平均相对误差、最大相对误差、最大绝对误差2115. **判定条件**:所有数据类型测试通过的状态212213其中,**精度标准**部分必须列出所有支持的数据类型及其对应的精度要求,确保报告的可读性和可追溯性。214215---216217## 7. 示例218219### 7.1 输入示例(Triton算子实现)220221```python222@triton.jit223def triton_abs(in_ptr0, out_ptr0, XBLOCK: tl.constexpr, XBLOCK_SUB: tl.constexpr):224 offset = tl.program_id(0) * XBLOCK225 base1 = tl.arange(0, XBLOCK_SUB)226 loops1: tl.constexpr = (XBLOCK + XBLOCK_SUB - 1) // XBLOCK_SUB227 for loop1 in range(loops1):228 x0_prime = offset + (loop1 * XBLOCK_SUB) + base1229 x0 = offset + (loop1 * XBLOCK_SUB) + base1230 tmp0 = tl.load(in_ptr0 + (x0), None)231 tmp2 = tl.abs(tmp0)232 tl.store(out_ptr0 + (x0), tmp2, None)233```234235### 7.2 输出示例(精度报告)236237```238================================================================================239 Triton算子精度验证报告 240--------------------------------------------------------------------------------241[验证配置]:242 数据类型: float32 (Single Precision)243 MERE阈值: 1.220703e-04244 MARE阈值: 1.220703e-03 (10×MERE阈值)245 小值域阈值: 1.000000e-07246--------------------------------------------------------------------------------247[验证结果]:248 验证结果: PASS249 样本总数: 4096250--------------------------------------------------------------------------------251[误差指标]:252 平均相对误差(MERE): 0.000000e+00253 阈值要求: MERE < 1.220703e-04254 最大相对误差(MARE): 0.000000e+00255 阈值要求: MARE < 1.220703e-03256--------------------------------------------------------------------------------257[判定条件]:258 ✓ MERE < 阈值: True259 ✓ MARE < 10×阈值: True260 ✓ 总体结果: True261================================================================================262```263264---265266## 8. 注意事项2672681. **环境配置**:确保已正确安装Triton、PyTorch和必要的依赖2692. **数据类型**:不同数据类型有不同的精度验证规则,请根据算子类型选择合适的数据类型2703. **参数配置**:根据算子复杂度和硬件资源调整测试参数(如ncore、xblock等)2714. **错误处理**:如果测试失败,检查Triton算子实现和Torch对比实现是否功能一致2725. **报告解读**:根据精度报告中的误差指标调整算子实现以提高精度273274---275276## 9. 故障处理277278| 问题 | 可能原因 | 解决方案 |279|-----|---------|---------|280| Triton kernel编译失败 | Triton语法错误或版本不兼容 | 检查Triton语法,确保Triton版本与代码兼容 |281| 精度验证失败 | 算子实现逻辑错误或精度损失 | 检查算子实现,调整算法以提高精度 |282| NPU设备不可用 | 未安装torch_npu或设备未正确配置 | 安装torch_npu,检查NPU设备状态 |283| 内存不足 | 测试数据过大 | 减小测试数据规模或调整参数配置 |284285---286287## 10. 参考资料288289- [Triton官方文档](https://triton-lang.org/documentation.html)290- [PyTorch官方文档](https://pytorch.org/docs/stable/index.html)291- [昇腾PyTorch插件文档](https://www.hiascend.com/document/detail/zh/pytorchplugin/master/ptplugin/ptplugin_000001.html)292293---