From 52a1ddc55c532d39285e8afcac54aadba4c29716 Mon Sep 17 00:00:00 2001 From: ElleElleWu <1608928702@qq.com> Date: Fri, 8 May 2026 19:16:43 +0800 Subject: [PATCH 1/9] =?UTF-8?q?[triton]=20=E6=9B=B4=E6=96=B0=E7=B2=BE?= =?UTF-8?q?=E5=BA=A6=E8=AF=84=E6=B5=8B=E6=A0=87=E5=87=86=EF=BC=8C=E6=94=AF?= =?UTF-8?q?=E6=8C=81=E5=8D=95=E6=A0=87=E6=9D=86=EF=BC=88MERE=20&=20MARE=20?= =?UTF-8?q?&=20=E5=B0=8F=E5=80=BC=E5=9F=9F=E6=A0=87=E5=87=86=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- skills/triton/kernel-verifier/SKILL.md | 54 +- .../triton/kernel-verifier/scripts/verify.py | 554 +++++++++--------- 2 files changed, 335 insertions(+), 273 deletions(-) diff --git a/skills/triton/kernel-verifier/SKILL.md b/skills/triton/kernel-verifier/SKILL.md index e1cf7552..bddb8758 100644 --- a/skills/triton/kernel-verifier/SKILL.md +++ b/skills/triton/kernel-verifier/SKILL.md @@ -327,21 +327,44 @@ benchmark.py 启动时按 `--triton_impl_name` 推导对应的 verify_result 文 ## 精度阈值说明 -验证使用基于数据类型的 **MERE/MARE 双门限相对误差**判定(NPU Benchmark 标准),与 `torch.allclose` 不同。 +验证使用基于数据类型的 **MERE/MARE 相对误差 + 小值域绝对误差** 双轨判定(NPU Benchmark 标准)。 -**判定公式**(必须同时满足): +### 常规精度标准(正常值域) + +对 `|golden| >= small_value_threshold` 的元素计算相对误差: ``` MERE < threshold 且 MARE < 10 × threshold ``` 其中: -- `MERE` = mean(|actual - golden| / max(|golden|, threshold)),平均相对误差 -- `MARE` = max(|actual - golden| / max(|golden|, threshold)),最大相对误差 +- `MERE` = mean(|actual - golden| / (|golden| + 1e-7)),平均相对误差 +- `MARE` = max(|actual - golden| / (|golden| + 1e-7)),最大相对误差 - 计算前两侧统一升 float32,避免低精度 dtype 自身误差污染 -- 分母用 `clamp(min=threshold)` 而非 `+epsilon`:当 `|golden| < threshold`(参考值已小到 dtype 精度极限)时,rel_err 退化为 `|diff| / threshold`,等价于按绝对误差归一化,避免零值/极小值附近误报 +- 分母用 `|golden| + 1e-7` 防止除零 + +### 小值域标准(极小值) + +当 golden 接近 0 时,相对误差计算不稳定,因此采用绝对误差判定: + +``` +ErrorCount = sum(I(|golden| < small_value_threshold 且 |actual - golden| > small_value_error)) +通过条件:ErrorCount <= 2 +``` + +### 三种场景的分支判定 -**dtype 阈值表**(2 的幂次方): +根据 golden 值的分布,分为三种场景: + +| 场景 | 判定条件 | 检查标准 | +|------|---------|---------| +| **全部在小值域** | 所有 `\|golden\| < small_value_threshold` | 仅检查小值域标准 | +| **全部在正常值域** | 所有 `\|golden\| >= small_value_threshold` | 仅检查常规精度标准(MERE/MARE) | +| **混合情况** | 同时存在小值和正常值 | 将输出**分割**成两个子集分别评估:小值子集检查小值域标准,正常子集检查常规精度标准,两部分都通过才算整体通过 | + +### 阈值表 + +**常规精度阈值**(用于相对误差判定): | 数据类型 | threshold | MERE 上限 | MARE 上限 (10×t) | |---------|-----------|-----------|------------------| @@ -353,12 +376,25 @@ MERE < threshold 且 MARE < 10 × threshold | `float8_e5m2` | 2⁻² = 0.25 | 0.25 | 2.5 | | 其他 dtype(fallback) | 2⁻¹³ | 1.22e-4 | 1.22e-3 | -**比对前置检查**(按顺序,任一失败即判 fail): +**小值域阈值表**: + +| 数据类型 | small_value_threshold | small_value_error | +|---------|----------------------|-------------------| +| `float16` | 2⁻¹¹ ≈ 4.88e-4 | 2⁻¹⁶ ≈ 1.53e-5 | +| `bfloat16` | 2⁻⁸ ≈ 3.91e-3 | 2⁻¹⁶ ≈ 1.53e-5 | +| `float32` | 2⁻¹⁴ ≈ 6.10e-5 | 2⁻³⁰ ≈ 9.31e-10 | +| `hifloat32` | 2⁻¹² ≈ 2.44e-4 | 2⁻²⁸ ≈ 3.73e-9 | +| `float8_e4m3` | 2⁻⁴ = 0.0625 | 2⁻⁶ = 0.015625 | +| `float8_e5m2` | 2⁻³ = 0.125 | 2⁻⁵ = 0.03125 | +| 其他 dtype(fallback) | 2⁻¹⁴ | 2⁻³⁰ | + +### 比对前置检查(按顺序,任一失败即判 fail) + 1. 形状必须一致 2. NaN 位置必须完全一致(mask 按位相等) 3. Inf 位置和符号必须完全一致 -4. `bool` dtype:要求 `torch.equal` 完全相等,不进入 MERE/MARE 判定 -5. 仅在 `finite_mask` 上做 MERE/MARE 计算;当 dtype 不一致时 impl 会被 cast 到 golden 的 dtype +4. `bool` dtype:要求 `torch.equal` 完全相等,不进入精度判定 +5. 仅在 `finite_mask` 上做精度计算;当 dtype 不一致时 impl 会被 cast 到 golden 的 dtype --- diff --git a/skills/triton/kernel-verifier/scripts/verify.py b/skills/triton/kernel-verifier/scripts/verify.py index 10ca3120..07f252b6 100644 --- a/skills/triton/kernel-verifier/scripts/verify.py +++ b/skills/triton/kernel-verifier/scripts/verify.py @@ -4,23 +4,6 @@ 多 shape 模式下:每个 shape 独立 try/except,全部跑完后落盘 verify_result.json。 策略 A:passed < total 即整体判失败(exit 1),同时失败清单记录在 JSON 的 `failures` 字段。 -精度判定标准: - allclose 样式逐元素判定: - abs(actual - golden) <= atol + rtol * abs(golden) - -当前阈值: - FLOAT32: - rtol = 1.220703125e-4 | 2**(-13) - atol = 1e-5 - - FLOAT16: - rtol = 9.765625e-4 | 2**(-10) - atol = 1e-3 - - BFLOAT16: - rtol = 7.8125e-3 | 2**(-7) - atol = 1e-2 - 用法: python verify.py --op_name <算子名> [--verify_dir <验证目录>] [--timeout <超时秒数>] """ @@ -79,74 +62,157 @@ def cleanup_npu_memory(): gc.collect() -def get_allclose_tolerance(data_type): - """根据数据类型获取 allclose 样式精度阈值。 +def get_limit(data_type): + """根据数据类型获取精度阈值 - 使用 2 的幂次方阈值(与 NPU Benchmark 标准一致) + 参考文档: 精度对比方法.md + 数据类型: FLOAT16, BFLOAT16, FLOAT32, HiFloat32, FLOAT8 E4M3, FLOAT8 E5M2 + 判定标准: MERE < threshold 且 MARE < 10 * threshold + + + 阈值表: + | 数据类型 | 阈值 (2^n) | 十进制值 | + |--------------|----------------|---------------| + | FLOAT16 | 2^{-10} | 0.0009765625 | + | BFLOAT16 | 2^{-7} | 0.0078125 | + | FLOAT32 | 2^{-13} | 0.0001220703 | + | HiFloat32 | 2^{-11} | 0.0004882812 | + | FLOAT8 E4M3 | 2^{-3} | 0.125 | + | FLOAT8 E5M2 | 2^{-2} | 0.25 | + + 由于 torch.dtype 中没有直接定义 HiFloat32,可通过字符串传入 "hifloat32" 获取对应阈值。 + """ # noqa: E501 + + + import torch + + # 支持字符串类型(用于 HiFloat32 或其他自定义类型) + + if isinstance(data_type, str): + str_to_threshold = { + "float16": 2**(-10), + "bfloat16": 2**(-7), + "float32": 2**(-13), + "hifloat32": 2**(-11), + "float8_e4m3": 2**(-3), + "float8_e5m2": 2**(-2), + "fp8_e4m3": 2**(-3), + "fp8_e5m2": 2**(-2), + } + return str_to_threshold.get(data_type.lower(), 2**(-13)) + + # torch.dtype 类型映射 + dtype_threshold_map = { + torch.float16: 2**(-10), # FLOAT16 + torch.bfloat16: 2**(-7), # BFLOAT16 + torch.float32: 2**(-13), # FLOAT32 + } + + # 安全获取 FP8 类型(PyTorch 2.0+ 支持) + # FLOAT8 E4M3: 2^{-3} + float8_e4m3 = getattr(torch, 'float8_e4m3fn', None) or getattr(torch, 'float8_e4m3', None) + if float8_e4m3 is not None: + dtype_threshold_map[float8_e4m3] = 2**(-3) + + # FLOAT8 E5M2: 2^{-2} + float8_e5m2 = getattr(torch, 'float8_e5m2fn', None) or getattr(torch, 'float8_e5m2', None) + if float8_e5m2 is not None: + dtype_threshold_map[float8_e5m2] = 2**(-2) + + return dtype_threshold_map.get(data_type, 2**(-13)) - 判定标准: - abs(actual - golden) <= atol + rtol * abs(golden) - 当前采用阈值: - FLOAT32: - rtol = 2^{-13} = 1.220703125e-4 - atol = 1e-5 +def get_small_value_threshold(data_type): + """获取小值域阈值 (Small Value Threshold)。 - FLOAT16: - rtol = 2^{-10} = 9.765625e-4 - atol = 1e-3 + 当 |golden| < threshold 时,采用小值域通过标准评估精度。 - BFLOAT16: - rtol = 2^{-7} = 7.8125e-3 - atol = 1e-2 + 阈值表: + | 数据类型 | 小值域阈值 (2^n) | 十进制值 | + |--------------|------------------|---------------| + | FLOAT16 | 2^{-11} | 4.8828125e-4 | + | BFLOAT16 | 2^{-8} | 0.00390625 | + | FLOAT32 | 2^{-14} | 6.1035156e-5 | + | HiFloat32 | 2^{-12} | 2.4414062e-4 | + | FLOAT8 E4M3 | 2^{-4} | 0.0625 | + | FLOAT8 E5M2 | 2^{-3} | 0.125 | """ import torch - default_tol = { - "rtol": 2**(-13), - "atol": 1e-5, + if isinstance(data_type, str): + str_to_threshold = { + "float16": 2**(-11), + "bfloat16": 2**(-8), + "float32": 2**(-14), + "hifloat32": 2**(-12), + "float8_e4m3": 2**(-4), + "float8_e5m2": 2**(-3), + "fp8_e4m3": 2**(-4), + "fp8_e5m2": 2**(-3), + } + return str_to_threshold.get(data_type.lower(), 2**(-14)) + + dtype_threshold_map = { + torch.float16: 2**(-11), + torch.bfloat16: 2**(-8), + torch.float32: 2**(-14), } + float8_e4m3 = getattr(torch, 'float8_e4m3fn', None) or getattr(torch, 'float8_e4m3', None) + if float8_e4m3 is not None: + dtype_threshold_map[float8_e4m3] = 2**(-4) + + float8_e5m2 = getattr(torch, 'float8_e5m2fn', None) or getattr(torch, 'float8_e5m2', None) + if float8_e5m2 is not None: + dtype_threshold_map[float8_e5m2] = 2**(-3) + + return dtype_threshold_map.get(data_type, 2**(-14)) + + +def get_small_value_error(data_type): + """获取小值域 error 指标。 + + 当 |golden| < small_value_threshold 时,若 |actual - golden| > error 则计为错误。 + + 阈值表: + | 数据类型 | 小值域 error (2^n) | 十进制值 | + |--------------|-------------------|---------------| + | FLOAT16 | 2^{-16} | 1.5258789e-5 | + | BFLOAT16 | 2^{-16} | 1.5258789e-5 | + | FLOAT32 | 2^{-30} | 9.3132257e-10 | + | HiFloat32 | 2^{-28} | 3.7252903e-9 | + | FLOAT8 E4M3 | 2^{-6} | 0.015625 | + | FLOAT8 E5M2 | 2^{-5} | 0.03125 | + """ + import torch + if isinstance(data_type, str): - key = data_type.lower().replace("torch.", "") - str_to_tol = { - "float32": { - "rtol": 2**(-13), - "atol": 1e-5, - }, - "float": { - "rtol": 2**(-13), - "atol": 1e-5, - }, - "float16": { - "rtol": 2**(-10), - "atol": 1e-3, - }, - "half": { - "rtol": 2**(-10), - "atol": 1e-3, - }, - "bfloat16": { - "rtol": 2**(-7), - "atol": 1e-2, - }, + str_to_error = { + "float16": 2**(-16), + "bfloat16": 2**(-16), + "float32": 2**(-30), + "hifloat32": 2**(-28), + "float8_e4m3": 2**(-6), + "float8_e5m2": 2**(-5), + "fp8_e4m3": 2**(-6), + "fp8_e5m2": 2**(-5), } - return str_to_tol.get(key, default_tol) - - dtype_to_tol = { - torch.float32: { - "rtol": 2**(-13), - "atol": 1e-5, - }, - torch.float16: { - "rtol": 2**(-10), - "atol": 1e-3, - }, - torch.bfloat16: { - "rtol": 2**(-7), - "atol": 1e-2, - }, + return str_to_error.get(data_type.lower(), 2**(-30)) + + dtype_error_map = { + torch.float16: 2**(-16), + torch.bfloat16: 2**(-16), + torch.float32: 2**(-30), } - return dtype_to_tol.get(data_type, default_tol) + float8_e4m3 = getattr(torch, 'float8_e4m3fn', None) or getattr(torch, 'float8_e4m3', None) + if float8_e4m3 is not None: + dtype_error_map[float8_e4m3] = 2**(-6) + + float8_e5m2 = getattr(torch, 'float8_e5m2fn', None) or getattr(torch, 'float8_e5m2', None) + if float8_e5m2 is not None: + dtype_error_map[float8_e5m2] = 2**(-5) + + return dtype_error_map.get(data_type, 2**(-30)) def resolve_input_provider(torch_module): @@ -157,11 +223,13 @@ def resolve_input_provider(torch_module): elif hasattr(torch_module, "get_inputs"): return [torch_module.get_inputs()], 1 else: - raise AttributeError("模块必须提供 get_inputs() 或 get_input_groups() 方法") + raise AttributeError( + f"模块必须提供 get_inputs() 或 get_input_groups() 方法" + ) def compare(fw_out, impl_out, data_type): - """对比框架输出和实现输出。""" + """对比框架输出和实现输出""" import torch fw_flat = fw_out.flatten().detach().cpu() @@ -173,6 +241,7 @@ def compare(fw_out, impl_out, data_type): impl_flat = torch.tensor(impl_flat, dtype=fw_flat.dtype) size = fw_flat.numel() + print(f" 总元素数: {size}", file=sys.stderr) if fw_flat.shape != impl_flat.shape: raise AssertionError( @@ -188,6 +257,9 @@ def compare(fw_out, impl_out, data_type): f"验证失败,NaN 位置不匹配: Framework={fw_nan_count}/{size}, " f"Implementation={impl_nan_count}/{size}" ) + fw_nan_count = fw_nan_mask.sum().item() + if fw_nan_count > 0: + print(f" NaN 检查通过: NaN数量={fw_nan_count}", file=sys.stderr) fw_inf_mask = torch.isinf(fw_flat) impl_inf_mask = torch.isinf(impl_flat) @@ -198,6 +270,9 @@ def compare(fw_out, impl_out, data_type): f"验证失败,Inf 位置不匹配: Framework={fw_inf_count}/{size}, " f"Implementation={impl_inf_count}/{size}" ) + fw_inf_count = fw_inf_mask.sum().item() + if fw_inf_count > 0: + print(f" Inf 检查通过: Inf数量={fw_inf_count}", file=sys.stderr) if fw_inf_mask.any(): if not torch.equal( @@ -210,136 +285,139 @@ def compare(fw_out, impl_out, data_type): finite_count = finite_mask.sum().item() if finite_count == 0: - print("警告: 所有值都是非有限值,跳过精度检查") + print(" 警告: 所有值都是非有限值,跳过精度检查", file=sys.stderr) return + print(f" 有限值数量: {finite_count}", file=sys.stderr) + fw_finite = fw_flat[finite_mask] impl_finite = impl_flat[finite_mask] if fw_finite.dtype == torch.bool: if not torch.equal(fw_finite, impl_finite): raise AssertionError(f"验证失败,布尔值不匹配: dtype={data_type}") + print(f" 布尔值检查通过", file=sys.stderr) return if impl_finite.dtype != fw_finite.dtype: impl_finite = impl_finite.to(fw_finite.dtype) + print(f" dtype转换: impl -> {fw_finite.dtype}", file=sys.stderr) + + # 执行 NPU Benchmark 精度验证 + _check_accuracy_npu_benchmark(fw_finite, impl_finite, data_type) - # 执行 allclose 精度验证 - _check_accuracy_allclose(fw_finite, impl_finite, data_type) +def _check_accuracy_npu_benchmark(golden, actual, data_type): + """执行 NPU Benchmark 精度验证(单标杆比对)。 -def _check_accuracy_allclose(golden, actual, data_type): - """执行 allclose 精度验证。 + 验证两个张量的数值一致性: + - 计算 MERE(平均相对误差)和 MARE(最大相对误差) + - 使用 2 的幂次方作为阈值 + - 判定标准: + - 若所有 golden 都落在小值域(|golden| < small_value_threshold),仅检查小值域通过标准 + - 若所有 golden 都不在小值域,检查 MERE < threshold 且 MARE < 10 * threshold + - 若混合情况,将输出分割成小值部分和正常部分分别评估: + - 小值部分:检查小值域通过标准 + - 正常部分:仅对非小值计算 MERE/MARE 并检查常规精度标准 + - 两部分都通过才算整体通过 - 判定标准: - abs(actual - golden) <= atol + rtol * abs(golden) + 小值域通过标准: + - ErrorCount = sum(I(|golden| < small_value_threshold and |actual - golden| > error)) + - 通过条件:ErrorCount <= 2 Args: - golden: 参考输出,通常是 PyTorch framework 输出 - actual: 被测实现输出,通常是 Triton-Ascend 输出 - data_type: 数据类型,用于获取对应阈值 + golden: 参考输出(金标准) + actual: 被测实现输出 + data_type: 数据类型,用于获取对应的阈值 Raises: AssertionError: 当精度验证未通过时 """ import torch + # 统一转换为 float32 进行计算 golden_f = golden.float() actual_f = actual.float() - if golden_f.shape != actual_f.shape: - raise AssertionError( - f"验证失败,输出形状不一致: golden={golden_f.shape}, actual={actual_f.shape}" - ) - - numel = golden_f.numel() - if numel == 0: - return - - tol = get_allclose_tolerance(data_type) - rtol = tol["rtol"] - atol = tol["atol"] - + threshold = get_limit(data_type) diff = (actual_f - golden_f).abs() - golden_abs = golden_f.abs() - - allowed_error = atol + rtol * golden_abs - close_mask = diff <= allowed_error - allclose_ok = bool(close_mask.all().item()) - - if not allclose_ok: - failed_close_mask = ~close_mask - failed_close_count = int(failed_close_mask.sum().item()) - pass_rate = 1.0 - failed_close_count / max(numel, 1) - - max_abs_err = diff.max().item() - mean_abs_err = diff.mean().item() - max_allowed_err = allowed_error.max().item() - mean_allowed_err = allowed_error.mean().item() - - # 为了日志可读,计算一个诊断用相对误差。 - # 注意:该 relative_error 只用于错误信息展示,不参与判定。 - rel_denom_floor = atol / rtol - rel_denom = torch.clamp(golden_abs, min=rel_denom_floor) - relative_error = diff / rel_denom - max_rel_err = relative_error.max().item() - mean_rel_err = relative_error.mean().item() - - failed_indices = torch.where(failed_close_mask)[0] - num_failed_to_show = min(10, len(failed_indices)) - - topk = min(10, numel) - top_rel_values, top_rel_indices = torch.topk(relative_error, k=topk) - - error_msg = ( - "验证失败,输出不一致:\n" - f" dtype={data_type}\n" - f" numel={numel}\n" - f" allclose_ok={allclose_ok}\n" - f" pass_rate={pass_rate:.6%}\n" - f" failed_close_count={failed_close_count}/{numel}\n" - "\n" - "阈值配置:\n" - f" rtol={rtol:.12e}\n" - f" atol={atol:.12e}\n" - f" rel_denom_floor=atol/rtol={rel_denom_floor:.12e} # 仅用于日志中的相对误差\n" - "\n" - "误差统计:\n" - f" max_abs_err={max_abs_err:.12e}\n" - f" mean_abs_err={mean_abs_err:.12e}\n" - f" max_rel_err={max_rel_err:.12e} # 仅日志\n" - f" mean_rel_err={mean_rel_err:.12e} # 仅日志\n" - f" max_allowed_err={max_allowed_err:.12e}\n" - f" mean_allowed_err={mean_allowed_err:.12e}\n" - ) - if failed_close_count > 0: - error_msg += f"\n前 {num_failed_to_show} 个 allclose 失败点:\n" - for i in range(num_failed_to_show): - idx = failed_indices[i].item() - error_msg += ( - f" 位置[{idx}]: " - f"golden={golden_f[idx].item():.12e}, " - f"actual={actual_f[idx].item():.12e}, " - f"abs_err={diff[idx].item():.12e}, " - f"allowed={allowed_error[idx].item():.12e}, " - f"rel_err={relative_error[idx].item():.12e}\n" - ) - - error_msg += f"\n相对误差最大的前 {topk} 个点,注意仅用于诊断,不参与判定:\n" - for i in range(topk): - idx = top_rel_indices[i].item() + # 小值域通过标准 + small_value_threshold = get_small_value_threshold(data_type) + small_value_error = get_small_value_error(data_type) + small_value_mask = golden_f.abs() < small_value_threshold + + # 判定标准: + # - 若所有 golden 都落在小值域,仅检查小值域通过标准 + # - 若所有 golden 都不在小值域,检查常规精度标准 + # - 若混合情况,将输出分割成小值部分和正常部分分别评估 + has_small_value = small_value_mask.any().item() + has_normal_value = (~small_value_mask).any().item() + + is_pass = True + normal_MERE = None + normal_MARE = None + + total_elements = golden_f.numel() + small_count = small_value_mask.sum().item() + normal_count = total_elements - small_count + + print(f" [精度检查] 总元素数={total_elements}, 小值域元素数={small_count}, " + f"正常值域元素数={normal_count}", file=sys.stderr) + + if has_small_value: + small_value_errors = diff[small_value_mask] + error_count = (small_value_errors > small_value_error).sum().item() + small_value_pass = error_count <= 2 + is_pass = is_pass and small_value_pass + print(f" [小值域检查] threshold={small_value_threshold:.6e}, " + f"error_limit={small_value_error:.6e}, ErrorCount={error_count}, " + f"通过={small_value_pass}", file=sys.stderr) + + if has_normal_value: + # 正常部分:仅对非小值计算相对误差 + normal_golden = golden_f[~small_value_mask] + normal_actual = actual_f[~small_value_mask] + normal_diff = diff[~small_value_mask] + normal_denom = normal_golden.abs() + 1e-7 + normal_relative_error = normal_diff / normal_denom + normal_MERE = normal_relative_error.mean().item() + normal_MARE = normal_relative_error.max().item() + normal_pass = (normal_MERE < threshold) and (normal_MARE < 10 * threshold) + is_pass = is_pass and normal_pass + print(f" [正常值域检查] MERE={normal_MERE:.6e}, MARE={normal_MARE:.6e}, " + f"threshold={threshold}, 通过={normal_pass}", file=sys.stderr) + + if not is_pass: + error_msg = f"验证失败,输出不一致: dtype={data_type}, threshold={threshold}\n" + + if has_small_value and not small_value_pass: error_msg += ( - f" 位置[{idx}]: " - f"golden={golden_f[idx].item():.12e}, " - f"actual={actual_f[idx].item():.12e}, " - f"abs_err={diff[idx].item():.12e}, " - f"allowed={allowed_error[idx].item():.12e}, " - f"rel_err={relative_error[idx].item():.12e}\n" + f"小值域未通过: small_value_threshold={small_value_threshold:.6e}, " + f"small_value_error={small_value_error:.6e}, ErrorCount={error_count}\n" ) + if has_normal_value and not normal_pass: + error_msg += ( + f"正常值域未通过: MERE={normal_MERE:.6e}, MARE={normal_MARE:.6e}\n" + ) + # 收集正常值域中超出阈值的样本 + mismatch_mask = normal_relative_error > threshold + mismatch_indices = torch.where(mismatch_mask)[0] + num_to_show = min(10, len(mismatch_indices)) + if len(mismatch_indices) > 0: + error_msg += f"前 {num_to_show} 个超出阈值的值:\n" + for i in range(num_to_show): + idx = mismatch_indices[i].item() + error_msg += ( + f" 位置[{idx}]: framework={normal_golden[idx]:.6e}, " + f"impl={normal_actual[idx]:.6e}, " + f"相对误差={normal_relative_error[idx]:.6e}\n" + ) raise AssertionError(error_msg) + print(f" [精度检查] 通过", file=sys.stderr) + def run_single_case( framework_model, @@ -347,12 +425,13 @@ def run_single_case( inputs, device, case_idx, - total_cases, + total_cases ): """验证单组输入。失败时抛出 AssertionError。""" import torch print(f" 测试第 {case_idx}/{total_cases} 组输入...", file=sys.stderr) + print(f" 输入描述: {describe_input(inputs)}", file=sys.stderr) inputs_for_impl = [ x.to(device) if isinstance(x, torch.Tensor) else x @@ -364,14 +443,18 @@ def run_single_case( ] with torch.no_grad(): - impl_output = impl_model(*inputs_for_impl) + print(f" 执行框架模型...", file=sys.stderr) framework_output = framework_model(*inputs_for_framework) + print(f" 执行实现模型...", file=sys.stderr) + impl_output = impl_model(*inputs_for_impl) if not isinstance(framework_output, (list, tuple)): framework_output = [framework_output] if not isinstance(impl_output, (list, tuple)): impl_output = [impl_output] + print(f" 输出数量: framework={len(framework_output)}, impl={len(impl_output)}", file=sys.stderr) + if len(framework_output) != len(impl_output): raise AssertionError( f"[用例 {case_idx}/{total_cases}] 输出数量不一致: " @@ -386,25 +469,17 @@ def run_single_case( ) if isinstance(fw_out, torch.Tensor) and isinstance(impl_out, torch.Tensor): + print(f" 比对输出 {i}: shape={list(fw_out.shape)}, dtype={fw_out.dtype}", file=sys.stderr) try: data_type = fw_out.dtype compare(fw_out, impl_out, data_type) except AssertionError as e: - raise AssertionError(f"[用例 {case_idx}/{total_cases}] 输出 {i}: {str(e)}") from e + raise AssertionError(f"[用例 {case_idx}/{total_cases}] {str(e)}") from e else: - if fw_out != impl_out: - raise AssertionError( - f"[用例 {case_idx}/{total_cases}] 输出 {i} 非 Tensor 值不一致: " - f"framework={fw_out}, impl={impl_out}" - ) - - -def verify_implementations( - op_name, - verify_dir, - triton_impl_name="triton_ascend_impl", - output_path=None, -): + print(f" 输出 {i} 非 Tensor,跳过精度比对", file=sys.stderr) + + +def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_impl", output_path=None): """验证框架实现和生成实现的结果一致性。 每个 shape 独立 try/except,全部跑完后写 verify_result.json。 @@ -435,7 +510,15 @@ def verify_implementations( failures = [] passed_cases = 0 + print(f"=" * 60, file=sys.stderr) + print(f"开始验证算子: {op_name}", file=sys.stderr) + print(f"总测试用例数: {total_cases}", file=sys.stderr) + print(f"=" * 60, file=sys.stderr) + for case_idx, inputs in enumerate(input_groups, start=1): + print(f"\n{'-' * 50}", file=sys.stderr) + print(f"[用例 {case_idx}/{total_cases}] 开始执行", file=sys.stderr) + input_desc = describe_input(inputs) framework_model = None impl_model = None @@ -452,21 +535,14 @@ def verify_implementations( impl_model = ModelNew(*init_params).to(device) run_single_case( - framework_model, - impl_model, - inputs, - device, - case_idx, - total_cases, + framework_model, impl_model, inputs, device, case_idx, total_cases ) passed_cases += 1 + print(f"[用例 {case_idx}/{total_cases}] 通过", file=sys.stderr) except Exception as e: err_detail = traceback.format_exc() - print( - f" [用例 {case_idx}/{total_cases}] 失败: {type(e).__name__}: {e}", - file=sys.stderr, - ) + print(f"[用例 {case_idx}/{total_cases}] 失败: {type(e).__name__}: {e}", file=sys.stderr) failures.append({ "case_idx": case_idx, "input_desc": input_desc, @@ -480,7 +556,11 @@ def verify_implementations( cleanup_npu_memory() failed_cases = total_cases - passed_cases + print(f"\n{'=' * 60}", file=sys.stderr) + print(f"验证完成: {passed_cases}/{total_cases} 通过, {failed_cases} 失败", file=sys.stderr) + print(f"{'=' * 60}", file=sys.stderr) + # 落盘 verify_result.json if output_path is None: output_path = os.path.join(verify_dir, "verify_result.json") @@ -515,29 +595,18 @@ def verify_implementations( parser = argparse.ArgumentParser(description="算子验证脚本") parser.add_argument("--op_name", required=True, help="算子名称") parser.add_argument( - "--verify_dir", - default=".", - help=( - "验证目录,包含 {op_name}_torch.py 和 " - "{op_name}_triton_ascend_impl.py(默认当前目录)" - ), + "--verify_dir", default=".", + help="验证目录,包含 {op_name}_torch.py 和 {op_name}_triton_ascend_impl.py(默认当前目录)", ) parser.add_argument("--timeout", type=int, default=900, help="超时秒数(默认 900)") parser.add_argument( - "--triton_impl_name", - default="triton_ascend_impl", + "--triton_impl_name", default="triton_ascend_impl", help="Triton 实现模块名(不含 op_name 前缀,默认 triton_ascend_impl)", ) parser.add_argument( - "--output", - default=None, + "--output", default=None, help="验证结果 JSON 输出路径(默认 {verify_dir}/verify_result.json)", ) - parser.add_argument( - "--_run", - action="store_true", - help=argparse.SUPPRESS, - ) args = parser.parse_args() @@ -546,57 +615,14 @@ def verify_implementations( print(f"错误: 验证目录不存在: {verify_dir}", file=sys.stderr) sys.exit(1) - if args._run: - # 子进程模式:直接执行验证逻辑 - try: - passed, total = verify_implementations( - args.op_name, - verify_dir, - args.triton_impl_name, - args.output, - ) - except Exception as e: - print(f"{e}", file=sys.stderr) - traceback.print_exc() - sys.exit(1) - - # 策略 A:passed < total → exit 1 - sys.exit(0 if passed == total and total > 0 else 1) - - else: - # 主进程模式:启动子进程执行验证,超时后 kill 子进程 - cmd = [ - sys.executable, - os.path.abspath(__file__), - "--op_name", - args.op_name, - "--verify_dir", - verify_dir, - "--triton_impl_name", - args.triton_impl_name, - "--_run", - ] - - if args.output: - cmd.extend(["--output", args.output]) - - try: - proc = subprocess.Popen( - cmd, - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - ) - stdout, stderr = proc.communicate(timeout=args.timeout) - - sys.stdout.buffer.write(stdout) - sys.stdout.buffer.flush() - sys.stderr.buffer.write(stderr) - sys.stderr.buffer.flush() - - sys.exit(proc.returncode) + try: + passed, total = verify_implementations( + args.op_name, verify_dir, args.triton_impl_name, args.output + ) + except Exception as e: + print(f"{e}", file=sys.stderr) + traceback.print_exc() + sys.exit(1) - except subprocess.TimeoutExpired: - proc.kill() - proc.wait() - print(f"验证超时({args.timeout}秒),已终止子进程", file=sys.stderr) - sys.exit(1) + # 策略 A:passed < total → exit 1 + sys.exit(0 if passed == total and total > 0 else 1) \ No newline at end of file From 4f668048d01407971b1d138b039957a08171e34e Mon Sep 17 00:00:00 2001 From: ElleElleWu <1608928702@qq.com> Date: Sat, 9 May 2026 17:18:36 +0800 Subject: [PATCH 2/9] =?UTF-8?q?[triton]=20=E6=94=B9=E5=9B=9Eallclose?= =?UTF-8?q?=E6=A0=B7=E5=BC=8F=EF=BC=8C=E4=BF=AE=E5=A4=8D=E6=89=B9=E8=B7=91?= =?UTF-8?q?session.md=E6=9C=AA=E7=94=9F=E6=88=90=E7=9A=84=E9=97=AE?= =?UTF-8?q?=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../triton/kernel-verifier/scripts/verify.py | 572 +++++++++--------- utils/run_benchmark_triton.sh | 49 +- 2 files changed, 312 insertions(+), 309 deletions(-) diff --git a/skills/triton/kernel-verifier/scripts/verify.py b/skills/triton/kernel-verifier/scripts/verify.py index 07f252b6..0130648e 100644 --- a/skills/triton/kernel-verifier/scripts/verify.py +++ b/skills/triton/kernel-verifier/scripts/verify.py @@ -4,6 +4,23 @@ 多 shape 模式下:每个 shape 独立 try/except,全部跑完后落盘 verify_result.json。 策略 A:passed < total 即整体判失败(exit 1),同时失败清单记录在 JSON 的 `failures` 字段。 +精度判定标准: + allclose 样式逐元素判定: + abs(actual - golden) <= atol + rtol * abs(golden) + +当前阈值: + FLOAT32: + rtol = 1.220703125e-4 | 2**(-13) + atol = 1e-5 + + FLOAT16: + rtol = 9.765625e-4 | 2**(-10) + atol = 1e-3 + + BFLOAT16: + rtol = 7.8125e-3 | 2**(-7) + atol = 1e-2 + 用法: python verify.py --op_name <算子名> [--verify_dir <验证目录>] [--timeout <超时秒数>] """ @@ -62,157 +79,74 @@ def cleanup_npu_memory(): gc.collect() -def get_limit(data_type): - """根据数据类型获取精度阈值 - 使用 2 的幂次方阈值(与 NPU Benchmark 标准一致) - 参考文档: 精度对比方法.md - 数据类型: FLOAT16, BFLOAT16, FLOAT32, HiFloat32, FLOAT8 E4M3, FLOAT8 E5M2 - 判定标准: MERE < threshold 且 MARE < 10 * threshold - - - 阈值表: - | 数据类型 | 阈值 (2^n) | 十进制值 | - |--------------|----------------|---------------| - | FLOAT16 | 2^{-10} | 0.0009765625 | - | BFLOAT16 | 2^{-7} | 0.0078125 | - | FLOAT32 | 2^{-13} | 0.0001220703 | - | HiFloat32 | 2^{-11} | 0.0004882812 | - | FLOAT8 E4M3 | 2^{-3} | 0.125 | - | FLOAT8 E5M2 | 2^{-2} | 0.25 | +def get_allclose_tolerance(data_type): + """根据数据类型获取 allclose 样式精度阈值。 + 判定标准: + abs(actual - golden) <= atol + rtol * abs(golden) - 由于 torch.dtype 中没有直接定义 HiFloat32,可通过字符串传入 "hifloat32" 获取对应阈值。 - """ # noqa: E501 + 当前采用阈值: + FLOAT32: + rtol = 2^{-13} = 1.220703125e-4 + atol = 1e-5 + FLOAT16: + rtol = 2^{-10} = 9.765625e-4 + atol = 1e-3 - import torch - - # 支持字符串类型(用于 HiFloat32 或其他自定义类型) - - if isinstance(data_type, str): - str_to_threshold = { - "float16": 2**(-10), - "bfloat16": 2**(-7), - "float32": 2**(-13), - "hifloat32": 2**(-11), - "float8_e4m3": 2**(-3), - "float8_e5m2": 2**(-2), - "fp8_e4m3": 2**(-3), - "fp8_e5m2": 2**(-2), - } - return str_to_threshold.get(data_type.lower(), 2**(-13)) + BFLOAT16: + rtol = 2^{-7} = 7.8125e-3 + atol = 1e-2 - # torch.dtype 类型映射 - dtype_threshold_map = { - torch.float16: 2**(-10), # FLOAT16 - torch.bfloat16: 2**(-7), # BFLOAT16 - torch.float32: 2**(-13), # FLOAT32 - } - - # 安全获取 FP8 类型(PyTorch 2.0+ 支持) - # FLOAT8 E4M3: 2^{-3} - float8_e4m3 = getattr(torch, 'float8_e4m3fn', None) or getattr(torch, 'float8_e4m3', None) - if float8_e4m3 is not None: - dtype_threshold_map[float8_e4m3] = 2**(-3) - - # FLOAT8 E5M2: 2^{-2} - float8_e5m2 = getattr(torch, 'float8_e5m2fn', None) or getattr(torch, 'float8_e5m2', None) - if float8_e5m2 is not None: - dtype_threshold_map[float8_e5m2] = 2**(-2) - - return dtype_threshold_map.get(data_type, 2**(-13)) - - -def get_small_value_threshold(data_type): - """获取小值域阈值 (Small Value Threshold)。 - - 当 |golden| < threshold 时,采用小值域通过标准评估精度。 - - 阈值表: - | 数据类型 | 小值域阈值 (2^n) | 十进制值 | - |--------------|------------------|---------------| - | FLOAT16 | 2^{-11} | 4.8828125e-4 | - | BFLOAT16 | 2^{-8} | 0.00390625 | - | FLOAT32 | 2^{-14} | 6.1035156e-5 | - | HiFloat32 | 2^{-12} | 2.4414062e-4 | - | FLOAT8 E4M3 | 2^{-4} | 0.0625 | - | FLOAT8 E5M2 | 2^{-3} | 0.125 | """ import torch - if isinstance(data_type, str): - str_to_threshold = { - "float16": 2**(-11), - "bfloat16": 2**(-8), - "float32": 2**(-14), - "hifloat32": 2**(-12), - "float8_e4m3": 2**(-4), - "float8_e5m2": 2**(-3), - "fp8_e4m3": 2**(-4), - "fp8_e5m2": 2**(-3), - } - return str_to_threshold.get(data_type.lower(), 2**(-14)) + default_tol = { + "rtol": 2**(-13), + "atol": 1e-5, - dtype_threshold_map = { - torch.float16: 2**(-11), - torch.bfloat16: 2**(-8), - torch.float32: 2**(-14), } - - float8_e4m3 = getattr(torch, 'float8_e4m3fn', None) or getattr(torch, 'float8_e4m3', None) - if float8_e4m3 is not None: - dtype_threshold_map[float8_e4m3] = 2**(-4) - - float8_e5m2 = getattr(torch, 'float8_e5m2fn', None) or getattr(torch, 'float8_e5m2', None) - if float8_e5m2 is not None: - dtype_threshold_map[float8_e5m2] = 2**(-3) - - return dtype_threshold_map.get(data_type, 2**(-14)) - - -def get_small_value_error(data_type): - """获取小值域 error 指标。 - - 当 |golden| < small_value_threshold 时,若 |actual - golden| > error 则计为错误。 - - 阈值表: - | 数据类型 | 小值域 error (2^n) | 十进制值 | - |--------------|-------------------|---------------| - | FLOAT16 | 2^{-16} | 1.5258789e-5 | - | BFLOAT16 | 2^{-16} | 1.5258789e-5 | - | FLOAT32 | 2^{-30} | 9.3132257e-10 | - | HiFloat32 | 2^{-28} | 3.7252903e-9 | - | FLOAT8 E4M3 | 2^{-6} | 0.015625 | - | FLOAT8 E5M2 | 2^{-5} | 0.03125 | - """ - import torch - if isinstance(data_type, str): - str_to_error = { - "float16": 2**(-16), - "bfloat16": 2**(-16), - "float32": 2**(-30), - "hifloat32": 2**(-28), - "float8_e4m3": 2**(-6), - "float8_e5m2": 2**(-5), - "fp8_e4m3": 2**(-6), - "fp8_e5m2": 2**(-5), + key = data_type.lower().replace("torch.", "") + str_to_tol = { + "float32": { + "rtol": 2**(-13), + "atol": 1e-5, + }, + "float": { + "rtol": 2**(-13), + "atol": 1e-5, + }, + "float16": { + "rtol": 2**(-10), + "atol": 1e-3, + }, + "half": { + "rtol": 2**(-10), + "atol": 1e-3, + }, + "bfloat16": { + "rtol": 2**(-7), + "atol": 1e-2, + }, } - return str_to_error.get(data_type.lower(), 2**(-30)) - - dtype_error_map = { - torch.float16: 2**(-16), - torch.bfloat16: 2**(-16), - torch.float32: 2**(-30), + return str_to_tol.get(key, default_tol) + + dtype_to_tol = { + torch.float32: { + "rtol": 2**(-13), + "atol": 1e-5, + }, + torch.float16: { + "rtol": 2**(-10), + "atol": 1e-3, + }, + torch.bfloat16: { + "rtol": 2**(-7), + "atol": 1e-2, + }, } - float8_e4m3 = getattr(torch, 'float8_e4m3fn', None) or getattr(torch, 'float8_e4m3', None) - if float8_e4m3 is not None: - dtype_error_map[float8_e4m3] = 2**(-6) - - float8_e5m2 = getattr(torch, 'float8_e5m2fn', None) or getattr(torch, 'float8_e5m2', None) - if float8_e5m2 is not None: - dtype_error_map[float8_e5m2] = 2**(-5) - - return dtype_error_map.get(data_type, 2**(-30)) + return dtype_to_tol.get(data_type, default_tol) def resolve_input_provider(torch_module): @@ -223,15 +157,11 @@ def resolve_input_provider(torch_module): elif hasattr(torch_module, "get_inputs"): return [torch_module.get_inputs()], 1 else: - raise AttributeError( - f"模块必须提供 get_inputs() 或 get_input_groups() 方法" - ) - + raise AttributeError("模块必须提供 get_inputs() 或 get_input_groups() 方法") def compare(fw_out, impl_out, data_type): - """对比框架输出和实现输出""" + """对比框架输出和实现输出。""" import torch - fw_flat = fw_out.flatten().detach().cpu() impl_flat = impl_out.flatten() @@ -241,7 +171,7 @@ def compare(fw_out, impl_out, data_type): impl_flat = torch.tensor(impl_flat, dtype=fw_flat.dtype) size = fw_flat.numel() - print(f" 总元素数: {size}", file=sys.stderr) + if fw_flat.shape != impl_flat.shape: raise AssertionError( @@ -257,10 +187,6 @@ def compare(fw_out, impl_out, data_type): f"验证失败,NaN 位置不匹配: Framework={fw_nan_count}/{size}, " f"Implementation={impl_nan_count}/{size}" ) - fw_nan_count = fw_nan_mask.sum().item() - if fw_nan_count > 0: - print(f" NaN 检查通过: NaN数量={fw_nan_count}", file=sys.stderr) - fw_inf_mask = torch.isinf(fw_flat) impl_inf_mask = torch.isinf(impl_flat) if not torch.equal(fw_inf_mask, impl_inf_mask): @@ -270,9 +196,6 @@ def compare(fw_out, impl_out, data_type): f"验证失败,Inf 位置不匹配: Framework={fw_inf_count}/{size}, " f"Implementation={impl_inf_count}/{size}" ) - fw_inf_count = fw_inf_mask.sum().item() - if fw_inf_count > 0: - print(f" Inf 检查通过: Inf数量={fw_inf_count}", file=sys.stderr) if fw_inf_mask.any(): if not torch.equal( @@ -285,10 +208,9 @@ def compare(fw_out, impl_out, data_type): finite_count = finite_mask.sum().item() if finite_count == 0: - print(" 警告: 所有值都是非有限值,跳过精度检查", file=sys.stderr) + print("警告: 所有值都是非有限值,跳过精度检查") return - print(f" 有限值数量: {finite_count}", file=sys.stderr) fw_finite = fw_flat[finite_mask] impl_finite = impl_flat[finite_mask] @@ -296,127 +218,132 @@ def compare(fw_out, impl_out, data_type): if fw_finite.dtype == torch.bool: if not torch.equal(fw_finite, impl_finite): raise AssertionError(f"验证失败,布尔值不匹配: dtype={data_type}") - print(f" 布尔值检查通过", file=sys.stderr) return if impl_finite.dtype != fw_finite.dtype: impl_finite = impl_finite.to(fw_finite.dtype) - print(f" dtype转换: impl -> {fw_finite.dtype}", file=sys.stderr) - # 执行 NPU Benchmark 精度验证 - _check_accuracy_npu_benchmark(fw_finite, impl_finite, data_type) + # 执行 allclose 精度验证 + _check_accuracy_allclose(fw_finite, impl_finite, data_type) + -def _check_accuracy_npu_benchmark(golden, actual, data_type): - """执行 NPU Benchmark 精度验证(单标杆比对)。 +def _check_accuracy_allclose(golden, actual, data_type): + """执行 allclose 精度验证。 - 验证两个张量的数值一致性: - - 计算 MERE(平均相对误差)和 MARE(最大相对误差) - - 使用 2 的幂次方作为阈值 - - 判定标准: - - 若所有 golden 都落在小值域(|golden| < small_value_threshold),仅检查小值域通过标准 - - 若所有 golden 都不在小值域,检查 MERE < threshold 且 MARE < 10 * threshold - - 若混合情况,将输出分割成小值部分和正常部分分别评估: - - 小值部分:检查小值域通过标准 - - 正常部分:仅对非小值计算 MERE/MARE 并检查常规精度标准 - - 两部分都通过才算整体通过 + 判定标准: + abs(actual - golden) <= atol + rtol * abs(golden) - 小值域通过标准: - - ErrorCount = sum(I(|golden| < small_value_threshold and |actual - golden| > error)) - - 通过条件:ErrorCount <= 2 Args: - golden: 参考输出(金标准) - actual: 被测实现输出 - data_type: 数据类型,用于获取对应的阈值 + golden: 参考输出,通常是 PyTorch framework 输出 + actual: 被测实现输出,通常是 Triton-Ascend 输出 + data_type: 数据类型,用于获取对应阈值 Raises: AssertionError: 当精度验证未通过时 """ import torch - # 统一转换为 float32 进行计算 + golden_f = golden.float() actual_f = actual.float() - threshold = get_limit(data_type) + if golden_f.shape != actual_f.shape: + raise AssertionError( + f"验证失败,输出形状不一致: golden={golden_f.shape}, actual={actual_f.shape}" + ) + + numel = golden_f.numel() + if numel == 0: + return + + tol = get_allclose_tolerance(data_type) + rtol = tol["rtol"] + atol = tol["atol"] + diff = (actual_f - golden_f).abs() + golden_abs = golden_f.abs() + + allowed_error = atol + rtol * golden_abs + close_mask = diff <= allowed_error + allclose_ok = bool(close_mask.all().item()) + + if not allclose_ok: + failed_close_mask = ~close_mask + failed_close_count = int(failed_close_mask.sum().item()) + pass_rate = 1.0 - failed_close_count / max(numel, 1) + + max_abs_err = diff.max().item() + mean_abs_err = diff.mean().item() + max_allowed_err = allowed_error.max().item() + mean_allowed_err = allowed_error.mean().item() + + # 为了日志可读,计算一个诊断用相对误差。 + # 注意:该 relative_error 只用于错误信息展示,不参与判定。 + rel_denom_floor = atol / rtol + rel_denom = torch.clamp(golden_abs, min=rel_denom_floor) + relative_error = diff / rel_denom + max_rel_err = relative_error.max().item() + mean_rel_err = relative_error.mean().item() + + failed_indices = torch.where(failed_close_mask)[0] + num_failed_to_show = min(10, len(failed_indices)) + + topk = min(10, numel) + top_rel_values, top_rel_indices = torch.topk(relative_error, k=topk) + + error_msg = ( + "验证失败,输出不一致:\n" + f" dtype={data_type}\n" + f" numel={numel}\n" + f" allclose_ok={allclose_ok}\n" + f" pass_rate={pass_rate:.6%}\n" + f" failed_close_count={failed_close_count}/{numel}\n" + "\n" + "阈值配置:\n" + f" rtol={rtol:.12e}\n" + f" atol={atol:.12e}\n" + f" rel_denom_floor=atol/rtol={rel_denom_floor:.12e} # 仅用于日志中的相对误差\n" + "\n" + "误差统计:\n" + f" max_abs_err={max_abs_err:.12e}\n" + f" mean_abs_err={mean_abs_err:.12e}\n" + f" max_rel_err={max_rel_err:.12e} # 仅日志\n" + f" mean_rel_err={mean_rel_err:.12e} # 仅日志\n" + f" max_allowed_err={max_allowed_err:.12e}\n" + f" mean_allowed_err={mean_allowed_err:.12e}\n" + ) - # 小值域通过标准 - small_value_threshold = get_small_value_threshold(data_type) - small_value_error = get_small_value_error(data_type) - small_value_mask = golden_f.abs() < small_value_threshold - - # 判定标准: - # - 若所有 golden 都落在小值域,仅检查小值域通过标准 - # - 若所有 golden 都不在小值域,检查常规精度标准 - # - 若混合情况,将输出分割成小值部分和正常部分分别评估 - has_small_value = small_value_mask.any().item() - has_normal_value = (~small_value_mask).any().item() - - is_pass = True - normal_MERE = None - normal_MARE = None - - total_elements = golden_f.numel() - small_count = small_value_mask.sum().item() - normal_count = total_elements - small_count - - print(f" [精度检查] 总元素数={total_elements}, 小值域元素数={small_count}, " - f"正常值域元素数={normal_count}", file=sys.stderr) - - if has_small_value: - small_value_errors = diff[small_value_mask] - error_count = (small_value_errors > small_value_error).sum().item() - small_value_pass = error_count <= 2 - is_pass = is_pass and small_value_pass - print(f" [小值域检查] threshold={small_value_threshold:.6e}, " - f"error_limit={small_value_error:.6e}, ErrorCount={error_count}, " - f"通过={small_value_pass}", file=sys.stderr) - - if has_normal_value: - # 正常部分:仅对非小值计算相对误差 - normal_golden = golden_f[~small_value_mask] - normal_actual = actual_f[~small_value_mask] - normal_diff = diff[~small_value_mask] - normal_denom = normal_golden.abs() + 1e-7 - normal_relative_error = normal_diff / normal_denom - normal_MERE = normal_relative_error.mean().item() - normal_MARE = normal_relative_error.max().item() - normal_pass = (normal_MERE < threshold) and (normal_MARE < 10 * threshold) - is_pass = is_pass and normal_pass - print(f" [正常值域检查] MERE={normal_MERE:.6e}, MARE={normal_MARE:.6e}, " - f"threshold={threshold}, 通过={normal_pass}", file=sys.stderr) - - if not is_pass: - error_msg = f"验证失败,输出不一致: dtype={data_type}, threshold={threshold}\n" - - if has_small_value and not small_value_pass: - error_msg += ( - f"小值域未通过: small_value_threshold={small_value_threshold:.6e}, " - f"small_value_error={small_value_error:.6e}, ErrorCount={error_count}\n" - ) + if failed_close_count > 0: + error_msg += f"\n前 {num_failed_to_show} 个 allclose 失败点:\n" + for i in range(num_failed_to_show): + idx = failed_indices[i].item() + error_msg += ( + f" 位置[{idx}]: " + f"golden={golden_f[idx].item():.12e}, " + f"actual={actual_f[idx].item():.12e}, " + f"abs_err={diff[idx].item():.12e}, " + f"allowed={allowed_error[idx].item():.12e}, " + f"rel_err={relative_error[idx].item():.12e}\n" + ) + + error_msg += f"\n相对误差最大的前 {topk} 个点,注意仅用于诊断,不参与判定:\n" + for i in range(topk): + idx = top_rel_indices[i].item() - if has_normal_value and not normal_pass: error_msg += ( - f"正常值域未通过: MERE={normal_MERE:.6e}, MARE={normal_MARE:.6e}\n" + f" 位置[{idx}]: " + f"golden={golden_f[idx].item():.12e}, " + f"actual={actual_f[idx].item():.12e}, " + f"abs_err={diff[idx].item():.12e}, " + f"allowed={allowed_error[idx].item():.12e}, " + f"rel_err={relative_error[idx].item():.12e}\n" ) - # 收集正常值域中超出阈值的样本 - mismatch_mask = normal_relative_error > threshold - mismatch_indices = torch.where(mismatch_mask)[0] - num_to_show = min(10, len(mismatch_indices)) - if len(mismatch_indices) > 0: - error_msg += f"前 {num_to_show} 个超出阈值的值:\n" - for i in range(num_to_show): - idx = mismatch_indices[i].item() - error_msg += ( - f" 位置[{idx}]: framework={normal_golden[idx]:.6e}, " - f"impl={normal_actual[idx]:.6e}, " - f"相对误差={normal_relative_error[idx]:.6e}\n" - ) + raise AssertionError(error_msg) - print(f" [精度检查] 通过", file=sys.stderr) + def run_single_case( @@ -425,13 +352,13 @@ def run_single_case( inputs, device, case_idx, - total_cases + total_cases, ): """验证单组输入。失败时抛出 AssertionError。""" import torch print(f" 测试第 {case_idx}/{total_cases} 组输入...", file=sys.stderr) - print(f" 输入描述: {describe_input(inputs)}", file=sys.stderr) + inputs_for_impl = [ x.to(device) if isinstance(x, torch.Tensor) else x @@ -443,17 +370,17 @@ def run_single_case( ] with torch.no_grad(): - print(f" 执行框架模型...", file=sys.stderr) - framework_output = framework_model(*inputs_for_framework) - print(f" 执行实现模型...", file=sys.stderr) impl_output = impl_model(*inputs_for_impl) + framework_output = framework_model(*inputs_for_framework) + + if not isinstance(framework_output, (list, tuple)): framework_output = [framework_output] if not isinstance(impl_output, (list, tuple)): impl_output = [impl_output] - print(f" 输出数量: framework={len(framework_output)}, impl={len(impl_output)}", file=sys.stderr) + if len(framework_output) != len(impl_output): raise AssertionError( @@ -469,17 +396,26 @@ def run_single_case( ) if isinstance(fw_out, torch.Tensor) and isinstance(impl_out, torch.Tensor): - print(f" 比对输出 {i}: shape={list(fw_out.shape)}, dtype={fw_out.dtype}", file=sys.stderr) + try: data_type = fw_out.dtype compare(fw_out, impl_out, data_type) except AssertionError as e: - raise AssertionError(f"[用例 {case_idx}/{total_cases}] {str(e)}") from e + raise AssertionError(f"[用例 {case_idx}/{total_cases}] 输出 {i}: {str(e)}") from e else: - print(f" 输出 {i} 非 Tensor,跳过精度比对", file=sys.stderr) - - -def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_impl", output_path=None): + if fw_out != impl_out: + raise AssertionError( + f"[用例 {case_idx}/{total_cases}] 输出 {i} 非 Tensor 值不一致: " + f"framework={fw_out}, impl={impl_out}" + ) + + +def verify_implementations( + op_name, + verify_dir, + triton_impl_name="triton_ascend_impl", + output_path=None, +): """验证框架实现和生成实现的结果一致性。 每个 shape 独立 try/except,全部跑完后写 verify_result.json。 @@ -510,14 +446,14 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ failures = [] passed_cases = 0 - print(f"=" * 60, file=sys.stderr) - print(f"开始验证算子: {op_name}", file=sys.stderr) - print(f"总测试用例数: {total_cases}", file=sys.stderr) - print(f"=" * 60, file=sys.stderr) + + + + for case_idx, inputs in enumerate(input_groups, start=1): - print(f"\n{'-' * 50}", file=sys.stderr) - print(f"[用例 {case_idx}/{total_cases}] 开始执行", file=sys.stderr) + + input_desc = describe_input(inputs) framework_model = None @@ -535,14 +471,22 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ impl_model = ModelNew(*init_params).to(device) run_single_case( - framework_model, impl_model, inputs, device, case_idx, total_cases + framework_model, + impl_model, + inputs, + device, + case_idx, + total_cases, ) passed_cases += 1 - print(f"[用例 {case_idx}/{total_cases}] 通过", file=sys.stderr) + except Exception as e: err_detail = traceback.format_exc() - print(f"[用例 {case_idx}/{total_cases}] 失败: {type(e).__name__}: {e}", file=sys.stderr) + print( + f" [用例 {case_idx}/{total_cases}] 失败: {type(e).__name__}: {e}", + file=sys.stderr, + ) failures.append({ "case_idx": case_idx, "input_desc": input_desc, @@ -556,11 +500,11 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ cleanup_npu_memory() failed_cases = total_cases - passed_cases - print(f"\n{'=' * 60}", file=sys.stderr) - print(f"验证完成: {passed_cases}/{total_cases} 通过, {failed_cases} 失败", file=sys.stderr) - print(f"{'=' * 60}", file=sys.stderr) - # 落盘 verify_result.json + + + + if output_path is None: output_path = os.path.join(verify_dir, "verify_result.json") @@ -595,18 +539,29 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ parser = argparse.ArgumentParser(description="算子验证脚本") parser.add_argument("--op_name", required=True, help="算子名称") parser.add_argument( - "--verify_dir", default=".", - help="验证目录,包含 {op_name}_torch.py 和 {op_name}_triton_ascend_impl.py(默认当前目录)", + "--verify_dir", + default=".", + help=( + "验证目录,包含 {op_name}_torch.py 和 " + "{op_name}_triton_ascend_impl.py(默认当前目录)" + ), ) parser.add_argument("--timeout", type=int, default=900, help="超时秒数(默认 900)") parser.add_argument( - "--triton_impl_name", default="triton_ascend_impl", + "--triton_impl_name", + default="triton_ascend_impl", help="Triton 实现模块名(不含 op_name 前缀,默认 triton_ascend_impl)", ) parser.add_argument( - "--output", default=None, + "--output", + default=None, help="验证结果 JSON 输出路径(默认 {verify_dir}/verify_result.json)", ) + parser.add_argument( + "--_run", + action="store_true", + help=argparse.SUPPRESS, + ) args = parser.parse_args() @@ -615,14 +570,57 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ print(f"错误: 验证目录不存在: {verify_dir}", file=sys.stderr) sys.exit(1) - try: - passed, total = verify_implementations( - args.op_name, verify_dir, args.triton_impl_name, args.output - ) - except Exception as e: - print(f"{e}", file=sys.stderr) - traceback.print_exc() - sys.exit(1) + if args._run: + # 子进程模式:直接执行验证逻辑 + try: + passed, total = verify_implementations( + args.op_name, + verify_dir, + args.triton_impl_name, + args.output, + ) + except Exception as e: + print(f"{e}", file=sys.stderr) + traceback.print_exc() + sys.exit(1) + + # 策略 A:passed < total → exit 1 + sys.exit(0 if passed == total and total > 0 else 1) + + else: + # 主进程模式:启动子进程执行验证,超时后 kill 子进程 + cmd = [ + sys.executable, + os.path.abspath(__file__), + "--op_name", + args.op_name, + "--verify_dir", + verify_dir, + "--triton_impl_name", + args.triton_impl_name, + "--_run", + ] + + if args.output: + cmd.extend(["--output", args.output]) + + try: + proc = subprocess.Popen( + cmd, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + ) + stdout, stderr = proc.communicate(timeout=args.timeout) + + sys.stdout.buffer.write(stdout) + sys.stdout.buffer.flush() + sys.stderr.buffer.write(stderr) + sys.stderr.buffer.flush() + + sys.exit(proc.returncode) - # 策略 A:passed < total → exit 1 - sys.exit(0 if passed == total and total > 0 else 1) \ No newline at end of file + except subprocess.TimeoutExpired: + proc.kill() + proc.wait() + print(f"验证超时({args.timeout}秒),已终止子进程", file=sys.stderr) + sys.exit(1) \ No newline at end of file diff --git a/utils/run_benchmark_triton.sh b/utils/run_benchmark_triton.sh index a1cd982f..95175846 100644 --- a/utils/run_benchmark_triton.sh +++ b/utils/run_benchmark_triton.sh @@ -210,6 +210,9 @@ if [[ "$USE_PARALLEL" == true ]]; then mkdir -p "$TARGET_OP_DIR" + # 预生成 session-id,调用后按 SID 精确取 jsonl,避免 ls -t 竞态 + SID=$(python3 -c 'import uuid;print(uuid.uuid4())') + START_TIME=$(date +%s) if [[ -f "$json_file" ]]; then @@ -219,6 +222,7 @@ if [[ "$USE_PARALLEL" == true ]]; then fi if claude -p "$PROMPT" \ + --session-id "$SID" \ --allowedTools 'Bash(*)' 'Read(*)' 'Write(*)' 'Edit(*)' 'Glob(*)' 'Grep(*)' 'Skill(*)' \ >> "${OUTPUT_DIR}/npu_${npu}.log" 2>&1; then @@ -251,19 +255,16 @@ if [[ "$USE_PARALLEL" == true ]]; then STATUS="fail" fi - # 串行重命名思维轨迹文件(带时间戳防止同名覆盖) - { - flock -x 201 - LATEST_JSONL=$(ls -t "$CLAUDE_PROJECT_DIR"/*.jsonl 2>/dev/null | head -1) - if [[ -n "$LATEST_JSONL" && -f "$LATEST_JSONL" ]]; then - BASENAME=$(basename "$LATEST_JSONL" .jsonl) - TIMESTAMP=$(date +%Y%m%d_%H%M%S) - mv "$LATEST_JSONL" "${CLAUDE_PROJECT_DIR}/${op_name}_${STATUS}_${TIMESTAMP}.jsonl" - if [[ -d "${CLAUDE_PROJECT_DIR}/${BASENAME}" ]]; then - mv "${CLAUDE_PROJECT_DIR}/${BASENAME}" "${CLAUDE_PROJECT_DIR}/${op_name}_${STATUS}_${TIMESTAMP}" - fi - fi - } 201>"${OUTPUT_DIR}/.trace_lock" + # 按 session-id 精确搬运思维轨迹(无需 flock,无竞态) + SRC_JSONL="${CLAUDE_PROJECT_DIR}/${SID}.jsonl" + if [[ -f "$SRC_JSONL" ]]; then + mv "$SRC_JSONL" "${TARGET_OP_DIR}/session.jsonl" + else + echo "[NPU $npu] ⚠ 未找到 session jsonl: ${SRC_JSONL}" >&2 + fi + if [[ -d "${CLAUDE_PROJECT_DIR}/${SID}" ]]; then + mv "${CLAUDE_PROJECT_DIR}/${SID}" "${TARGET_OP_DIR}/session_dir" + fi done # ========== Worker 进程结束 ========== ) & @@ -310,6 +311,9 @@ else START_TIME=$(date +%s) + # 预生成 session-id,调用后按 SID 精确取 jsonl + SID=$(python3 -c 'import uuid;print(uuid.uuid4())') + if [[ -f "$json_file" ]]; then PROMPT="生成一个基于 Triton-Ascend 框架的算子,参考${file}和${json_file}。目标设备架构为${ARCH},使用NPU=${NPU_ID},请将生成的代码文件输出至${TARGET_OP_DIR}/目录下。" else @@ -317,6 +321,7 @@ else fi if claude -p "$PROMPT" \ + --session-id "$SID" \ --allowedTools 'Bash(*)' 'Read(*)' 'Write(*)' 'Edit(*)' 'Glob(*)' 'Grep(*)' 'Skill(*)'; then END_TIME=$(date +%s) ELAPSED=$((END_TIME - START_TIME)) @@ -333,15 +338,15 @@ else STATUS="fail" fi - # 重命名思维轨迹文件(带时间戳防止同名覆盖) - LATEST_JSONL=$(ls -t "$CLAUDE_PROJECT_DIR"/*.jsonl 2>/dev/null | head -1) - if [[ -n "$LATEST_JSONL" && -f "$LATEST_JSONL" ]]; then - BASENAME=$(basename "$LATEST_JSONL" .jsonl) - TIMESTAMP=$(date +%Y%m%d_%H%M%S) - mv "$LATEST_JSONL" "${CLAUDE_PROJECT_DIR}/${op_name}_${STATUS}_${TIMESTAMP}.jsonl" - if [[ -d "${CLAUDE_PROJECT_DIR}/${BASENAME}" ]]; then - mv "${CLAUDE_PROJECT_DIR}/${BASENAME}" "${CLAUDE_PROJECT_DIR}/${op_name}_${STATUS}_${TIMESTAMP}" - fi + # 按 session-id 精确搬运思维轨迹 + SRC_JSONL="${CLAUDE_PROJECT_DIR}/${SID}.jsonl" + if [[ -f "$SRC_JSONL" ]]; then + mv "$SRC_JSONL" "${TARGET_OP_DIR}/session.jsonl" + else + echo "[NPU ${NPU_ID}] ⚠ 未找到 session jsonl: ${SRC_JSONL}" + fi + if [[ -d "${CLAUDE_PROJECT_DIR}/${SID}" ]]; then + mv "${CLAUDE_PROJECT_DIR}/${SID}" "${TARGET_OP_DIR}/session_dir" fi done fi From 41d5f76c6f515ed7912ec37ae8cbf10d984c5555 Mon Sep 17 00:00:00 2001 From: ElleElleWu <1608928702@qq.com> Date: Tue, 12 May 2026 10:46:05 +0800 Subject: [PATCH 3/9] [triton] verify change to MARE&MERE --- .../triton/kernel-verifier/scripts/verify.py | 368 +++++------------- 1 file changed, 106 insertions(+), 262 deletions(-) diff --git a/skills/triton/kernel-verifier/scripts/verify.py b/skills/triton/kernel-verifier/scripts/verify.py index 0130648e..a67d3162 100644 --- a/skills/triton/kernel-verifier/scripts/verify.py +++ b/skills/triton/kernel-verifier/scripts/verify.py @@ -4,23 +4,6 @@ 多 shape 模式下:每个 shape 独立 try/except,全部跑完后落盘 verify_result.json。 策略 A:passed < total 即整体判失败(exit 1),同时失败清单记录在 JSON 的 `failures` 字段。 -精度判定标准: - allclose 样式逐元素判定: - abs(actual - golden) <= atol + rtol * abs(golden) - -当前阈值: - FLOAT32: - rtol = 1.220703125e-4 | 2**(-13) - atol = 1e-5 - - FLOAT16: - rtol = 9.765625e-4 | 2**(-10) - atol = 1e-3 - - BFLOAT16: - rtol = 7.8125e-3 | 2**(-7) - atol = 1e-2 - 用法: python verify.py --op_name <算子名> [--verify_dir <验证目录>] [--timeout <超时秒数>] """ @@ -51,7 +34,6 @@ def describe_input(inputs): import torch except Exception: torch = None - descs = [] for x in inputs: if torch is not None and isinstance(x, torch.Tensor): @@ -79,74 +61,60 @@ def cleanup_npu_memory(): gc.collect() -def get_allclose_tolerance(data_type): - """根据数据类型获取 allclose 样式精度阈值。 - 判定标准: - abs(actual - golden) <= atol + rtol * abs(golden) +def get_limit(data_type): + """根据数据类型获取精度阈值 - 使用 2 的幂次方阈值(与 NPU Benchmark 标准一致) - 当前采用阈值: - FLOAT32: - rtol = 2^{-13} = 1.220703125e-4 - atol = 1e-5 + 参考文档: 精度对比方法.md + 数据类型: FLOAT16, BFLOAT16, FLOAT32, HiFloat32, FLOAT8 E4M3, FLOAT8 E5M2 + 判定标准: MERE < threshold 且 MARE < 10 * threshold - FLOAT16: - rtol = 2^{-10} = 9.765625e-4 - atol = 1e-3 + 阈值表: + | 数据类型 | 阈值 (2^n) | 十进制值 | + |--------------|----------------|---------------| + | FLOAT16 | 2^{-10} | 0.0009765625 | + | BFLOAT16 | 2^{-7} | 0.0078125 | + | FLOAT32 | 2^{-13} | 0.0001220703 | + | HiFloat32 | 2^{-11} | 0.0004882812 | + | FLOAT8 E4M3 | 2^{-3} | 0.125 | + | FLOAT8 E5M2 | 2^{-2} | 0.25 | - BFLOAT16: - rtol = 2^{-7} = 7.8125e-3 - atol = 1e-2 - - """ + 由于 torch.dtype 中没有直接定义 HiFloat32,可通过字符串传入 "hifloat32" 获取对应阈值。 + """ # noqa: E501 import torch - default_tol = { - "rtol": 2**(-13), - "atol": 1e-5, - - } + # 支持字符串类型(用于 HiFloat32 或其他自定义类型) if isinstance(data_type, str): - key = data_type.lower().replace("torch.", "") - str_to_tol = { - "float32": { - "rtol": 2**(-13), - "atol": 1e-5, - }, - "float": { - "rtol": 2**(-13), - "atol": 1e-5, - }, - "float16": { - "rtol": 2**(-10), - "atol": 1e-3, - }, - "half": { - "rtol": 2**(-10), - "atol": 1e-3, - }, - "bfloat16": { - "rtol": 2**(-7), - "atol": 1e-2, - }, + str_to_threshold = { + "float16": 2**(-10), + "bfloat16": 2**(-7), + "float32": 2**(-13), + "hifloat32": 2**(-11), + "float8_e4m3": 2**(-3), + "float8_e5m2": 2**(-2), + "fp8_e4m3": 2**(-3), + "fp8_e5m2": 2**(-2), } - return str_to_tol.get(key, default_tol) - - dtype_to_tol = { - torch.float32: { - "rtol": 2**(-13), - "atol": 1e-5, - }, - torch.float16: { - "rtol": 2**(-10), - "atol": 1e-3, - }, - torch.bfloat16: { - "rtol": 2**(-7), - "atol": 1e-2, - }, + return str_to_threshold.get(data_type.lower(), 2**(-13)) + + # torch.dtype 类型映射 + dtype_threshold_map = { + torch.float16: 2**(-10), # FLOAT16 + torch.bfloat16: 2**(-7), # BFLOAT16 + torch.float32: 2**(-13), # FLOAT32 } - return dtype_to_tol.get(data_type, default_tol) + # 安全获取 FP8 类型(PyTorch 2.0+ 支持) + # FLOAT8 E4M3: 2^{-3} + float8_e4m3 = getattr(torch, 'float8_e4m3fn', None) or getattr(torch, 'float8_e4m3', None) + if float8_e4m3 is not None: + dtype_threshold_map[float8_e4m3] = 2**(-3) + + # FLOAT8 E5M2: 2^{-2} + float8_e5m2 = getattr(torch, 'float8_e5m2fn', None) or getattr(torch, 'float8_e5m2', None) + if float8_e5m2 is not None: + dtype_threshold_map[float8_e5m2] = 2**(-2) + + return dtype_threshold_map.get(data_type, 2**(-13)) def resolve_input_provider(torch_module): @@ -157,14 +125,16 @@ def resolve_input_provider(torch_module): elif hasattr(torch_module, "get_inputs"): return [torch_module.get_inputs()], 1 else: - raise AttributeError("模块必须提供 get_inputs() 或 get_input_groups() 方法") + raise AttributeError( + f"模块必须提供 get_inputs() 或 get_input_groups() 方法" + ) + def compare(fw_out, impl_out, data_type): - """对比框架输出和实现输出。""" + """对比框架输出和实现输出""" import torch fw_flat = fw_out.flatten().detach().cpu() impl_flat = impl_out.flatten() - if isinstance(impl_flat, torch.Tensor): impl_flat = impl_flat.detach().cpu() else: @@ -172,7 +142,6 @@ def compare(fw_out, impl_out, data_type): size = fw_flat.numel() - if fw_flat.shape != impl_flat.shape: raise AssertionError( f"验证失败,输出形状不一致: framework={fw_flat.shape}, impl={impl_flat.shape}" @@ -187,6 +156,7 @@ def compare(fw_out, impl_out, data_type): f"验证失败,NaN 位置不匹配: Framework={fw_nan_count}/{size}, " f"Implementation={impl_nan_count}/{size}" ) + fw_inf_mask = torch.isinf(fw_flat) impl_inf_mask = torch.isinf(impl_flat) if not torch.equal(fw_inf_mask, impl_inf_mask): @@ -196,7 +166,6 @@ def compare(fw_out, impl_out, data_type): f"验证失败,Inf 位置不匹配: Framework={fw_inf_count}/{size}, " f"Implementation={impl_inf_count}/{size}" ) - if fw_inf_mask.any(): if not torch.equal( torch.sign(fw_flat[fw_inf_mask]), @@ -206,12 +175,10 @@ def compare(fw_out, impl_out, data_type): finite_mask = torch.isfinite(fw_flat) & torch.isfinite(impl_flat) finite_count = finite_mask.sum().item() - if finite_count == 0: print("警告: 所有值都是非有限值,跳过精度检查") return - fw_finite = fw_flat[finite_mask] impl_finite = impl_flat[finite_mask] @@ -223,143 +190,83 @@ def compare(fw_out, impl_out, data_type): if impl_finite.dtype != fw_finite.dtype: impl_finite = impl_finite.to(fw_finite.dtype) - # 执行 allclose 精度验证 - _check_accuracy_allclose(fw_finite, impl_finite, data_type) - - + # 执行 NPU Benchmark 精度验证 + _check_accuracy_npu_benchmark(fw_finite, impl_finite, data_type) -def _check_accuracy_allclose(golden, actual, data_type): - """执行 allclose 精度验证。 - 判定标准: - abs(actual - golden) <= atol + rtol * abs(golden) +def _check_accuracy_npu_benchmark(golden, actual, data_type): + """执行 NPU Benchmark 精度验证。 + 根据精度对比方法文档,验证两个张量的数值一致性: + - 计算 MERE(平均相对误差)和 MARE(最大相对误差) + - 使用 2 的幂次方作为阈值 + - 判定标准:MERE < threshold 且 MARE < 10 * threshold Args: - golden: 参考输出,通常是 PyTorch framework 输出 - actual: 被测实现输出,通常是 Triton-Ascend 输出 - data_type: 数据类型,用于获取对应阈值 + golden: 参考输出(金标准) + actual: 被测实现输出 + data_type: 数据类型,用于获取对应的阈值 Raises: AssertionError: 当精度验证未通过时 """ import torch - + # 统一转换为 float32 进行计算 golden_f = golden.float() actual_f = actual.float() - if golden_f.shape != actual_f.shape: - raise AssertionError( - f"验证失败,输出形状不一致: golden={golden_f.shape}, actual={actual_f.shape}" - ) - - numel = golden_f.numel() - if numel == 0: - return - - tol = get_allclose_tolerance(data_type) - rtol = tol["rtol"] - atol = tol["atol"] + # 先取 dtype 阈值,用作分母下界 clamp。 + # 当 |y_ref| < threshold 时,按 |diff| / threshold 衡量,等价于 + # "参考值已小到 dtype 精度极限时,改用绝对误差归一化",避免零值/极小值附近误报。 + threshold = get_limit(data_type) diff = (actual_f - golden_f).abs() - golden_abs = golden_f.abs() + denom = golden_f.abs().clamp(min=threshold) + relative_error = diff / denom - allowed_error = atol + rtol * golden_abs - close_mask = diff <= allowed_error - allclose_ok = bool(close_mask.all().item()) + # 计算误差指标 + MERE = relative_error.mean().item() # 平均相对误差 + MARE = relative_error.max().item() # 最大相对误差 - if not allclose_ok: - failed_close_mask = ~close_mask - failed_close_count = int(failed_close_mask.sum().item()) - pass_rate = 1.0 - failed_close_count / max(numel, 1) + # 判定标准:MERE < t 且 MARE < 10t + is_pass = (MERE < threshold) and (MARE < 10 * threshold) - max_abs_err = diff.max().item() - mean_abs_err = diff.mean().item() - max_allowed_err = allowed_error.max().item() - mean_allowed_err = allowed_error.mean().item() - - # 为了日志可读,计算一个诊断用相对误差。 - # 注意:该 relative_error 只用于错误信息展示,不参与判定。 - rel_denom_floor = atol / rtol - rel_denom = torch.clamp(golden_abs, min=rel_denom_floor) - relative_error = diff / rel_denom - max_rel_err = relative_error.max().item() - mean_rel_err = relative_error.mean().item() - - failed_indices = torch.where(failed_close_mask)[0] - num_failed_to_show = min(10, len(failed_indices)) - - topk = min(10, numel) - top_rel_values, top_rel_indices = torch.topk(relative_error, k=topk) + if not is_pass: + # 收集错误信息 + mismatch_mask = relative_error > threshold + mismatch_indices = torch.where(mismatch_mask)[0] + num_to_show = min(10, len(mismatch_indices)) error_msg = ( - "验证失败,输出不一致:\n" - f" dtype={data_type}\n" - f" numel={numel}\n" - f" allclose_ok={allclose_ok}\n" - f" pass_rate={pass_rate:.6%}\n" - f" failed_close_count={failed_close_count}/{numel}\n" - "\n" - "阈值配置:\n" - f" rtol={rtol:.12e}\n" - f" atol={atol:.12e}\n" - f" rel_denom_floor=atol/rtol={rel_denom_floor:.12e} # 仅用于日志中的相对误差\n" - "\n" - "误差统计:\n" - f" max_abs_err={max_abs_err:.12e}\n" - f" mean_abs_err={mean_abs_err:.12e}\n" - f" max_rel_err={max_rel_err:.12e} # 仅日志\n" - f" mean_rel_err={mean_rel_err:.12e} # 仅日志\n" - f" max_allowed_err={max_allowed_err:.12e}\n" - f" mean_allowed_err={mean_allowed_err:.12e}\n" + f"验证失败,输出不一致: MERE={MERE:.6e}, MARE={MARE:.6e}, " + f"dtype={data_type}, threshold={threshold}\n" ) - - if failed_close_count > 0: - error_msg += f"\n前 {num_failed_to_show} 个 allclose 失败点:\n" - for i in range(num_failed_to_show): - idx = failed_indices[i].item() + if len(mismatch_indices) > 0: + error_msg += f"前 {num_to_show} 个超出阈值的值:\n" + for i in range(num_to_show): + idx = mismatch_indices[i].item() error_msg += ( - f" 位置[{idx}]: " - f"golden={golden_f[idx].item():.12e}, " - f"actual={actual_f[idx].item():.12e}, " - f"abs_err={diff[idx].item():.12e}, " - f"allowed={allowed_error[idx].item():.12e}, " - f"rel_err={relative_error[idx].item():.12e}\n" + f" 位置[{idx}]: framework={golden[idx]:.6e}, " + f"impl={actual[idx]:.6e}, " + f"相对误差={relative_error[idx]:.6e}\n" ) - - error_msg += f"\n相对误差最大的前 {topk} 个点,注意仅用于诊断,不参与判定:\n" - for i in range(topk): - idx = top_rel_indices[i].item() - - error_msg += ( - f" 位置[{idx}]: " - f"golden={golden_f[idx].item():.12e}, " - f"actual={actual_f[idx].item():.12e}, " - f"abs_err={diff[idx].item():.12e}, " - f"allowed={allowed_error[idx].item():.12e}, " - f"rel_err={relative_error[idx].item():.12e}\n" - ) - raise AssertionError(error_msg) - - def run_single_case( framework_model, impl_model, inputs, device, case_idx, - total_cases, + total_cases ): """验证单组输入。失败时抛出 AssertionError。""" import torch print(f" 测试第 {case_idx}/{total_cases} 组输入...", file=sys.stderr) - inputs_for_impl = [ x.to(device) if isinstance(x, torch.Tensor) else x for x in inputs @@ -373,15 +280,11 @@ def run_single_case( impl_output = impl_model(*inputs_for_impl) framework_output = framework_model(*inputs_for_framework) - - if not isinstance(framework_output, (list, tuple)): framework_output = [framework_output] if not isinstance(impl_output, (list, tuple)): impl_output = [impl_output] - - if len(framework_output) != len(impl_output): raise AssertionError( f"[用例 {case_idx}/{total_cases}] 输出数量不一致: " @@ -394,28 +297,15 @@ def run_single_case( f"[用例 {case_idx}/{total_cases}] 输出 {i} 为 None: " f"framework={fw_out is None}, impl={impl_out is None}" ) - if isinstance(fw_out, torch.Tensor) and isinstance(impl_out, torch.Tensor): - try: data_type = fw_out.dtype compare(fw_out, impl_out, data_type) except AssertionError as e: - raise AssertionError(f"[用例 {case_idx}/{total_cases}] 输出 {i}: {str(e)}") from e - else: - if fw_out != impl_out: - raise AssertionError( - f"[用例 {case_idx}/{total_cases}] 输出 {i} 非 Tensor 值不一致: " - f"framework={fw_out}, impl={impl_out}" - ) + raise AssertionError(f"[用例 {case_idx}/{total_cases}] {str(e)}") from e -def verify_implementations( - op_name, - verify_dir, - triton_impl_name="triton_ascend_impl", - output_path=None, -): +def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_impl", output_path=None): """验证框架实现和生成实现的结果一致性。 每个 shape 独立 try/except,全部跑完后写 verify_result.json。 @@ -446,22 +336,12 @@ def verify_implementations( failures = [] passed_cases = 0 - - - - - for case_idx, inputs in enumerate(input_groups, start=1): - - - input_desc = describe_input(inputs) framework_model = None impl_model = None - try: init_params = get_init_inputs() - torch.manual_seed(0) torch.npu.manual_seed(0) framework_model = FrameworkModel(*init_params).to(device) @@ -471,29 +351,18 @@ def verify_implementations( impl_model = ModelNew(*init_params).to(device) run_single_case( - framework_model, - impl_model, - inputs, - device, - case_idx, - total_cases, + framework_model, impl_model, inputs, device, case_idx, total_cases ) passed_cases += 1 - - except Exception as e: err_detail = traceback.format_exc() - print( - f" [用例 {case_idx}/{total_cases}] 失败: {type(e).__name__}: {e}", - file=sys.stderr, - ) + print(f" [用例 {case_idx}/{total_cases}] 失败: {type(e).__name__}: {e}", file=sys.stderr) failures.append({ "case_idx": case_idx, "input_desc": input_desc, "error_type": type(e).__name__, "error_msg": truncate_error(err_detail), }) - finally: del framework_model del impl_model @@ -501,13 +370,9 @@ def verify_implementations( failed_cases = total_cases - passed_cases - - - - + # 落盘 verify_result.json if output_path is None: output_path = os.path.join(verify_dir, "verify_result.json") - result = { "op_name": op_name, "total_cases": total_cases, @@ -515,7 +380,6 @@ def verify_implementations( "failed_cases": failed_cases, "failures": failures, } - try: with open(output_path, "w", encoding="utf-8") as f: json.dump(result, f, indent=2, ensure_ascii=False) @@ -539,30 +403,22 @@ def verify_implementations( parser = argparse.ArgumentParser(description="算子验证脚本") parser.add_argument("--op_name", required=True, help="算子名称") parser.add_argument( - "--verify_dir", - default=".", - help=( - "验证目录,包含 {op_name}_torch.py 和 " - "{op_name}_triton_ascend_impl.py(默认当前目录)" - ), + "--verify_dir", default=".", + help="验证目录,包含 {op_name}_torch.py 和 {op_name}_triton_ascend_impl.py(默认当前目录)", ) parser.add_argument("--timeout", type=int, default=900, help="超时秒数(默认 900)") parser.add_argument( - "--triton_impl_name", - default="triton_ascend_impl", + "--triton_impl_name", default="triton_ascend_impl", help="Triton 实现模块名(不含 op_name 前缀,默认 triton_ascend_impl)", ) parser.add_argument( - "--output", - default=None, + "--output", default=None, help="验证结果 JSON 输出路径(默认 {verify_dir}/verify_result.json)", ) parser.add_argument( - "--_run", - action="store_true", - help=argparse.SUPPRESS, + "--_run", action="store_true", + help=argparse.SUPPRESS, # 内部参数:子进程模式,直接执行验证 ) - args = parser.parse_args() verify_dir = os.path.abspath(args.verify_dir) @@ -574,36 +430,25 @@ def verify_implementations( # 子进程模式:直接执行验证逻辑 try: passed, total = verify_implementations( - args.op_name, - verify_dir, - args.triton_impl_name, - args.output, + args.op_name, verify_dir, args.triton_impl_name, args.output ) except Exception as e: print(f"{e}", file=sys.stderr) traceback.print_exc() sys.exit(1) - # 策略 A:passed < total → exit 1 sys.exit(0 if passed == total and total > 0 else 1) - else: - # 主进程模式:启动子进程执行验证,超时后 kill 子进程 + # 主进程模式:启动子进程执行验证,超时后 kill 整个进程树 cmd = [ - sys.executable, - os.path.abspath(__file__), - "--op_name", - args.op_name, - "--verify_dir", - verify_dir, - "--triton_impl_name", - args.triton_impl_name, + sys.executable, os.path.abspath(__file__), + "--op_name", args.op_name, + "--verify_dir", verify_dir, + "--triton_impl_name", args.triton_impl_name, "--_run", ] - if args.output: cmd.extend(["--output", args.output]) - try: proc = subprocess.Popen( cmd, @@ -616,7 +461,6 @@ def verify_implementations( sys.stdout.buffer.flush() sys.stderr.buffer.write(stderr) sys.stderr.buffer.flush() - sys.exit(proc.returncode) except subprocess.TimeoutExpired: From 66db940596b71d53de120fedde882431d3a72ddc Mon Sep 17 00:00:00 2001 From: ElleElleWu <1608928702@qq.com> Date: Tue, 12 May 2026 10:59:38 +0800 Subject: [PATCH 4/9] =?UTF-8?q?=20=20[triton]=20=E7=B2=BE=E5=BA=A6?= =?UTF-8?q?=E8=AF=84=E6=B5=8B=E6=94=B9=E4=B8=BA=20max=5Ferror=5Fcap=20+=20?= =?UTF-8?q?matched=5Fratio=20+=20MERE=20=E4=B8=89=E9=A1=B9=E5=88=A4?= =?UTF-8?q?=E5=AE=9A?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- skills/triton/kernel-verifier/SKILL.md | 101 ++++---- .../triton/kernel-verifier/scripts/verify.py | 226 ++++++++++++------ 2 files changed, 202 insertions(+), 125 deletions(-) diff --git a/skills/triton/kernel-verifier/SKILL.md b/skills/triton/kernel-verifier/SKILL.md index bddb8758..34a677f4 100644 --- a/skills/triton/kernel-verifier/SKILL.md +++ b/skills/triton/kernel-verifier/SKILL.md @@ -327,66 +327,71 @@ benchmark.py 启动时按 `--triton_impl_name` 推导对应的 verify_result 文 ## 精度阈值说明 -验证使用基于数据类型的 **MERE/MARE 相对误差 + 小值域绝对误差** 双轨判定(NPU Benchmark 标准)。 +验证采用基于数据类型的 **元素级分类 matched + 三项整体判定**(NPU Benchmark 标准)。 -### 常规精度标准(正常值域) +### 元素级 matched 定义(分类) -对 `|golden| >= small_value_threshold` 的元素计算相对误差: +对每个 finite 元素 `i`,按 `|golden[i]|` 落入的类别分别判定: -``` -MERE < threshold 且 MARE < 10 × threshold -``` +- **小值域** `|golden[i]| < small_value_threshold`: + `matched[i] = (|actual[i] - golden[i]| <= small_value_error)` +- **正常域** `|golden[i]| >= small_value_threshold`: + `matched[i] = (|actual[i] - golden[i]| / (|golden[i]| + 1e-7) <= rel_threshold)` -其中: -- `MERE` = mean(|actual - golden| / (|golden| + 1e-7)),平均相对误差 -- `MARE` = max(|actual - golden| / (|golden| + 1e-7)),最大相对误差 -- 计算前两侧统一升 float32,避免低精度 dtype 自身误差污染 -- 分母用 `|golden| + 1e-7` 防止除零 +> 计算前两侧统一升 float32,避免低精度 dtype 自身误差污染。 +> 分母 `+1e-7` 仅为保险——正常域里 `|golden| >= small_value_threshold ≫ 1e-7`。 -### 小值域标准(极小值) +### 通过条件(三项 AND,全部满足才算通过) -当 golden 接近 0 时,相对误差计算不稳定,因此采用绝对误差判定: +1. **`max_error_cap`**:`max(|actual - golden|) <= 0.1`(全局绝对误差上限,dtype 无关) +2. **`required_matched_ratio`**:`sum(matched) / total_finite >= 0.98` +3. **`MERE`**:对**所有 finite 元素**计算 `rel_err = |diff| / (|golden| + 1e-7)` 再取均值,要求 `MERE < rel_threshold`。当 `total_finite == 0` 时本项自动通过。 -``` -ErrorCount = sum(I(|golden| < small_value_threshold 且 |actual - golden| > small_value_error)) -通过条件:ErrorCount <= 2 -``` +### 阈值表 -### 三种场景的分支判定 +| 数据类型 | small_value_threshold | small_value_error | rel_threshold (= MERE 上限) | +|---|---|---|---| +| `float16` | 2⁻¹¹ ≈ 4.88e-4 | 2⁻¹⁶ ≈ 1.53e-5 | 2⁻¹⁰ ≈ 9.77e-4 | +| `bfloat16` | 2⁻⁸ ≈ 3.91e-3 | 2⁻¹⁶ ≈ 1.53e-5 | 2⁻⁷ ≈ 7.81e-3 | +| `float32` | 2⁻¹⁴ ≈ 6.10e-5 | 2⁻³⁰ ≈ 9.31e-10 | 2⁻¹³ ≈ 1.22e-4 | +| `hifloat32` | 2⁻¹² ≈ 2.44e-4 | 2⁻²⁸ ≈ 3.73e-9 | 2⁻¹¹ ≈ 4.88e-4 | +| `float8_e4m3` | 2⁻⁴ = 0.0625 | 2⁻⁶ ≈ 0.015625 | 2⁻³ = 0.125 | +| `float8_e5m2` | 2⁻³ = 0.125 | 2⁻⁵ = 0.03125 | 2⁻² = 0.25 | +| 其他 dtype(fallback) | 2⁻¹⁴ | 2⁻³⁰ | 2⁻¹³ | -根据 golden 值的分布,分为三种场景: +### 失败时的 JSON 输出 -| 场景 | 判定条件 | 检查标准 | -|------|---------|---------| -| **全部在小值域** | 所有 `\|golden\| < small_value_threshold` | 仅检查小值域标准 | -| **全部在正常值域** | 所有 `\|golden\| >= small_value_threshold` | 仅检查常规精度标准(MERE/MARE) | -| **混合情况** | 同时存在小值和正常值 | 将输出**分割**成两个子集分别评估:小值子集检查小值域标准,正常子集检查常规精度标准,两部分都通过才算整体通过 | +`failures[*]` 在精度不达标时会带上结构化 `metrics`: -### 阈值表 +```json +{ + "case_idx": 1, + "input_desc": [...], + "error_type": "AccuracyError", + "error_msg": "...", + "metrics": { + "matched_ratio": 0.95, + "max_abs_diff": 0.2, + "MERE": 2.0e-4, + "rel_threshold": 1.22e-4, + "small_value_threshold": 6.10e-5, + "small_value_error": 9.31e-10, + "max_error_cap": 0.1, + "required_matched_ratio": 0.98, + "total_finite": 1000, + "matched_count": 950, + "small_count": 0, + "normal_count": 1000, + "checks": { + "max_error_cap": false, + "required_matched_ratio": false, + "MERE": false + } + } +} +``` -**常规精度阈值**(用于相对误差判定): - -| 数据类型 | threshold | MERE 上限 | MARE 上限 (10×t) | -|---------|-----------|-----------|------------------| -| `float16` | 2⁻¹⁰ ≈ 9.77e-4 | 9.77e-4 | 9.77e-3 | -| `bfloat16` | 2⁻⁷ ≈ 7.81e-3 | 7.81e-3 | 7.81e-2 | -| `float32` | 2⁻¹³ ≈ 1.22e-4 | 1.22e-4 | 1.22e-3 | -| `hifloat32` | 2⁻¹¹ ≈ 4.88e-4 | 4.88e-4 | 4.88e-3 | -| `float8_e4m3` | 2⁻³ = 0.125 | 0.125 | 1.25 | -| `float8_e5m2` | 2⁻² = 0.25 | 0.25 | 2.5 | -| 其他 dtype(fallback) | 2⁻¹³ | 1.22e-4 | 1.22e-3 | - -**小值域阈值表**: - -| 数据类型 | small_value_threshold | small_value_error | -|---------|----------------------|-------------------| -| `float16` | 2⁻¹¹ ≈ 4.88e-4 | 2⁻¹⁶ ≈ 1.53e-5 | -| `bfloat16` | 2⁻⁸ ≈ 3.91e-3 | 2⁻¹⁶ ≈ 1.53e-5 | -| `float32` | 2⁻¹⁴ ≈ 6.10e-5 | 2⁻³⁰ ≈ 9.31e-10 | -| `hifloat32` | 2⁻¹² ≈ 2.44e-4 | 2⁻²⁸ ≈ 3.73e-9 | -| `float8_e4m3` | 2⁻⁴ = 0.0625 | 2⁻⁶ = 0.015625 | -| `float8_e5m2` | 2⁻³ = 0.125 | 2⁻⁵ = 0.03125 | -| 其他 dtype(fallback) | 2⁻¹⁴ | 2⁻³⁰ | +`checks` 三个布尔位标记每项判定是否独立通过,下游可直接据此分类失败原因(绝对误差爆点 / 离群点过多 / 平均误差偏大)。 ### 比对前置检查(按顺序,任一失败即判 fail) diff --git a/skills/triton/kernel-verifier/scripts/verify.py b/skills/triton/kernel-verifier/scripts/verify.py index a67d3162..cd7c47bf 100644 --- a/skills/triton/kernel-verifier/scripts/verify.py +++ b/skills/triton/kernel-verifier/scripts/verify.py @@ -18,6 +18,18 @@ ERROR_MSG_LIMIT = 2000 +# 精度判定常量(dtype 无关) +MAX_ERROR_CAP = 0.1 +REQUIRED_MATCHED_RATIO = 0.98 + + +class AccuracyError(AssertionError): + """精度判定失败异常,附带结构化 metrics 便于下游统计。""" + + def __init__(self, message, metrics): + super().__init__(message) + self.metrics = metrics + def truncate_error(msg: str, limit: int = ERROR_MSG_LIMIT) -> str: if msg is None: @@ -61,60 +73,58 @@ def cleanup_npu_memory(): gc.collect() -def get_limit(data_type): - """根据数据类型获取精度阈值 - 使用 2 的幂次方阈值(与 NPU Benchmark 标准一致) +def get_limits(data_type): + """根据数据类型返回精度判定的三元组 (small_value_threshold, small_value_error, rel_threshold)。 - 参考文档: 精度对比方法.md - 数据类型: FLOAT16, BFLOAT16, FLOAT32, HiFloat32, FLOAT8 E4M3, FLOAT8 E5M2 - 判定标准: MERE < threshold 且 MARE < 10 * threshold + 参考 NPU Benchmark 精度对比方法: + - small_value_threshold:判定元素是否落在"小值域"的阈值 + - small_value_error:小值域元素的绝对误差上限 + - rel_threshold:正常值域元素的相对误差上限,同时也是 MERE 的判定阈值 - 阈值表: - | 数据类型 | 阈值 (2^n) | 十进制值 | - |--------------|----------------|---------------| - | FLOAT16 | 2^{-10} | 0.0009765625 | - | BFLOAT16 | 2^{-7} | 0.0078125 | - | FLOAT32 | 2^{-13} | 0.0001220703 | - | HiFloat32 | 2^{-11} | 0.0004882812 | - | FLOAT8 E4M3 | 2^{-3} | 0.125 | - | FLOAT8 E5M2 | 2^{-2} | 0.25 | + 阈值表: + | 数据类型 | small_value_threshold | small_value_error | rel_threshold | + |--------------|-----------------------|-------------------|---------------| + | FLOAT16 | 2^{-11} | 2^{-16} | 2^{-10} | + | BFLOAT16 | 2^{-8} | 2^{-16} | 2^{-7} | + | FLOAT32 | 2^{-14} | 2^{-30} | 2^{-13} | + | HiFloat32 | 2^{-12} | 2^{-28} | 2^{-11} | + | FLOAT8 E4M3 | 2^{-4} | 2^{-6} | 2^{-3} | + | FLOAT8 E5M2 | 2^{-3} | 2^{-5} | 2^{-2} | 由于 torch.dtype 中没有直接定义 HiFloat32,可通过字符串传入 "hifloat32" 获取对应阈值。 """ # noqa: E501 import torch - # 支持字符串类型(用于 HiFloat32 或其他自定义类型) + # 字符串映射(用于 HiFloat32 或其他自定义类型) + str_to_limits = { + "float16": (2**(-11), 2**(-16), 2**(-10)), + "bfloat16": (2**(-8), 2**(-16), 2**(-7)), + "float32": (2**(-14), 2**(-30), 2**(-13)), + "hifloat32": (2**(-12), 2**(-28), 2**(-11)), + "float8_e4m3": (2**(-4), 2**(-6), 2**(-3)), + "float8_e5m2": (2**(-3), 2**(-5), 2**(-2)), + "fp8_e4m3": (2**(-4), 2**(-6), 2**(-3)), + "fp8_e5m2": (2**(-3), 2**(-5), 2**(-2)), + } if isinstance(data_type, str): - str_to_threshold = { - "float16": 2**(-10), - "bfloat16": 2**(-7), - "float32": 2**(-13), - "hifloat32": 2**(-11), - "float8_e4m3": 2**(-3), - "float8_e5m2": 2**(-2), - "fp8_e4m3": 2**(-3), - "fp8_e5m2": 2**(-2), - } - return str_to_threshold.get(data_type.lower(), 2**(-13)) - - # torch.dtype 类型映射 - dtype_threshold_map = { - torch.float16: 2**(-10), # FLOAT16 - torch.bfloat16: 2**(-7), # BFLOAT16 - torch.float32: 2**(-13), # FLOAT32 + return str_to_limits.get(data_type.lower(), (2**(-14), 2**(-30), 2**(-13))) + + # torch.dtype 映射 + dtype_limits_map = { + torch.float16: (2**(-11), 2**(-16), 2**(-10)), + torch.bfloat16: (2**(-8), 2**(-16), 2**(-7)), + torch.float32: (2**(-14), 2**(-30), 2**(-13)), } - # 安全获取 FP8 类型(PyTorch 2.0+ 支持) - # FLOAT8 E4M3: 2^{-3} float8_e4m3 = getattr(torch, 'float8_e4m3fn', None) or getattr(torch, 'float8_e4m3', None) if float8_e4m3 is not None: - dtype_threshold_map[float8_e4m3] = 2**(-3) + dtype_limits_map[float8_e4m3] = (2**(-4), 2**(-6), 2**(-3)) - # FLOAT8 E5M2: 2^{-2} float8_e5m2 = getattr(torch, 'float8_e5m2fn', None) or getattr(torch, 'float8_e5m2', None) if float8_e5m2 is not None: - dtype_threshold_map[float8_e5m2] = 2**(-2) + dtype_limits_map[float8_e5m2] = (2**(-3), 2**(-5), 2**(-2)) - return dtype_threshold_map.get(data_type, 2**(-13)) + return dtype_limits_map.get(data_type, (2**(-14), 2**(-30), 2**(-13))) def resolve_input_provider(torch_module): @@ -195,63 +205,118 @@ def compare(fw_out, impl_out, data_type): def _check_accuracy_npu_benchmark(golden, actual, data_type): - """执行 NPU Benchmark 精度验证。 + """执行 NPU Benchmark 精度验证(分类 + 三项判定)。 + + 元素级 matched 定义: + - |golden| < small_value_threshold(小值域):|diff| <= small_value_error + - 否则(正常值域):|diff| / (|golden| + 1e-7) <= rel_threshold - 根据精度对比方法文档,验证两个张量的数值一致性: - - 计算 MERE(平均相对误差)和 MARE(最大相对误差) - - 使用 2 的幂次方作为阈值 - - 判定标准:MERE < threshold 且 MARE < 10 * threshold + 通过条件(三项 AND): + 1. max(|diff|) <= MAX_ERROR_CAP(0.1,dtype 无关的绝对误差上限) + 2. matched_ratio = sum(matched) / total_finite >= REQUIRED_MATCHED_RATIO(0.98) + 3. MERE < rel_threshold(对所有 finite 元素计算相对误差再取均值, + 分母统一用 |golden| + 1e-7 防除零) Args: golden: 参考输出(金标准) actual: 被测实现输出 - data_type: 数据类型,用于获取对应的阈值 + data_type: 数据类型,用于获取对应的阈值三元组 Raises: - AssertionError: 当精度验证未通过时 + AccuracyError: 当精度验证未通过时,异常的 metrics 属性携带结构化指标 """ import torch - # 统一转换为 float32 进行计算 + # 统一升 float32,避免低精度 dtype 自身误差污染计算 golden_f = golden.float() actual_f = actual.float() - # 先取 dtype 阈值,用作分母下界 clamp。 - # 当 |y_ref| < threshold 时,按 |diff| / threshold 衡量,等价于 - # "参考值已小到 dtype 精度极限时,改用绝对误差归一化",避免零值/极小值附近误报。 - threshold = get_limit(data_type) + sv_thr, sv_err, rel_thr = get_limits(data_type) - diff = (actual_f - golden_f).abs() - denom = golden_f.abs().clamp(min=threshold) - relative_error = diff / denom + abs_diff = (actual_f - golden_f).abs() + abs_golden = golden_f.abs() - # 计算误差指标 - MERE = relative_error.mean().item() # 平均相对误差 - MARE = relative_error.max().item() # 最大相对误差 + # 分桶 + small_mask = abs_golden < sv_thr + normal_mask = ~small_mask - # 判定标准:MERE < t 且 MARE < 10t - is_pass = (MERE < threshold) and (MARE < 10 * threshold) + # 元素级 matched + small_ok = abs_diff <= sv_err + rel_err = abs_diff / (abs_golden + 1e-7) + normal_ok = rel_err <= rel_thr + matched_mask = torch.where(small_mask, small_ok, normal_ok) - if not is_pass: - # 收集错误信息 - mismatch_mask = relative_error > threshold - mismatch_indices = torch.where(mismatch_mask)[0] - num_to_show = min(10, len(mismatch_indices)) + total_finite = matched_mask.numel() + matched_count = int(matched_mask.sum().item()) + matched_ratio = matched_count / total_finite if total_finite > 0 else 1.0 + max_abs_diff = abs_diff.max().item() if total_finite > 0 else 0.0 - error_msg = ( - f"验证失败,输出不一致: MERE={MERE:.6e}, MARE={MARE:.6e}, " - f"dtype={data_type}, threshold={threshold}\n" - ) - if len(mismatch_indices) > 0: - error_msg += f"前 {num_to_show} 个超出阈值的值:\n" - for i in range(num_to_show): - idx = mismatch_indices[i].item() + # MERE:对所有 finite 元素计算相对误差再取均值(分母统一 |golden| + 1e-7 防除零) + normal_count = int(normal_mask.sum().item()) + if total_finite > 0: + MERE = rel_err.mean().item() + mere_ok = MERE < rel_thr + else: + MERE = None + mere_ok = True + + cap_ok = max_abs_diff <= MAX_ERROR_CAP + ratio_ok = matched_ratio >= REQUIRED_MATCHED_RATIO + is_pass = cap_ok and ratio_ok and mere_ok + + if is_pass: + return + + metrics = { + "matched_ratio": matched_ratio, + "max_abs_diff": max_abs_diff, + "MERE": MERE, + "rel_threshold": rel_thr, + "small_value_threshold": sv_thr, + "small_value_error": sv_err, + "max_error_cap": MAX_ERROR_CAP, + "required_matched_ratio": REQUIRED_MATCHED_RATIO, + "total_finite": total_finite, + "matched_count": matched_count, + "small_count": int(small_mask.sum().item()), + "normal_count": normal_count, + "checks": { + "max_error_cap": cap_ok, + "required_matched_ratio": ratio_ok, + "MERE": mere_ok, + }, + } + + # 失败摘要 + 前 N 个 unmatched 位置(按所属桶注明判定标准) + unmatched_mask = ~matched_mask + unmatched_indices = torch.where(unmatched_mask)[0] + num_to_show = min(10, len(unmatched_indices)) + + mere_str = f"{MERE:.6e}" if MERE is not None else "n/a" + error_msg = ( + f"验证失败 dtype={data_type}: " + f"max_abs_diff={max_abs_diff:.6e} (cap={MAX_ERROR_CAP}, ok={cap_ok}), " + f"matched_ratio={matched_ratio:.6f} (req>={REQUIRED_MATCHED_RATIO}, ok={ratio_ok}), " + f"MERE={mere_str} (rel_thr={rel_thr:.6e}, ok={mere_ok}); " + f"small_count={metrics['small_count']}, normal_count={normal_count}\n" + ) + if num_to_show > 0: + error_msg += f"前 {num_to_show} 个未通过的位置:\n" + for i in range(num_to_show): + idx = unmatched_indices[i].item() + if small_mask[idx].item(): + error_msg += ( + f" 位置[{idx}] (小值域): framework={golden[idx]:.6e}, " + f"impl={actual[idx]:.6e}, |diff|={abs_diff[idx]:.6e} " + f"(允许<={sv_err:.6e})\n" + ) + else: error_msg += ( - f" 位置[{idx}]: framework={golden[idx]:.6e}, " - f"impl={actual[idx]:.6e}, " - f"相对误差={relative_error[idx]:.6e}\n" + f" 位置[{idx}] (正常域): framework={golden[idx]:.6e}, " + f"impl={actual[idx]:.6e}, 相对误差={rel_err[idx]:.6e} " + f"(允许<={rel_thr:.6e})\n" ) - raise AssertionError(error_msg) + raise AccuracyError(error_msg, metrics) def run_single_case( @@ -301,6 +366,10 @@ def run_single_case( try: data_type = fw_out.dtype compare(fw_out, impl_out, data_type) + except AccuracyError as e: + raise AccuracyError( + f"[用例 {case_idx}/{total_cases}] {str(e)}", e.metrics + ) from e except AssertionError as e: raise AssertionError(f"[用例 {case_idx}/{total_cases}] {str(e)}") from e @@ -357,12 +426,15 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ except Exception as e: err_detail = traceback.format_exc() print(f" [用例 {case_idx}/{total_cases}] 失败: {type(e).__name__}: {e}", file=sys.stderr) - failures.append({ + failure_entry = { "case_idx": case_idx, "input_desc": input_desc, "error_type": type(e).__name__, "error_msg": truncate_error(err_detail), - }) + } + if isinstance(e, AccuracyError): + failure_entry["metrics"] = e.metrics + failures.append(failure_entry) finally: del framework_model del impl_model From 5c2cd8b0e30328045605be32fddd8864f40fae1c Mon Sep 17 00:00:00 2001 From: w00934874 Date: Wed, 13 May 2026 14:58:31 +0800 Subject: [PATCH 5/9] =?UTF-8?q?[triton]=20required=5Fmatched=5Fratio?= =?UTF-8?q?=E8=B0=83=E6=95=B4=E8=87=B30.9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- skills/triton/kernel-verifier/SKILL.md | 2 +- .../triton/kernel-verifier/scripts/verify.py | 60 ++++--------------- 2 files changed, 12 insertions(+), 50 deletions(-) diff --git a/skills/triton/kernel-verifier/SKILL.md b/skills/triton/kernel-verifier/SKILL.md index 34a677f4..5b9e593b 100644 --- a/skills/triton/kernel-verifier/SKILL.md +++ b/skills/triton/kernel-verifier/SKILL.md @@ -344,7 +344,7 @@ benchmark.py 启动时按 `--triton_impl_name` 推导对应的 verify_result 文 ### 通过条件(三项 AND,全部满足才算通过) 1. **`max_error_cap`**:`max(|actual - golden|) <= 0.1`(全局绝对误差上限,dtype 无关) -2. **`required_matched_ratio`**:`sum(matched) / total_finite >= 0.98` +2. **`required_matched_ratio`**:`sum(matched) / total_finite >= 0.9` 3. **`MERE`**:对**所有 finite 元素**计算 `rel_err = |diff| / (|golden| + 1e-7)` 再取均值,要求 `MERE < rel_threshold`。当 `total_finite == 0` 时本项自动通过。 ### 阈值表 diff --git a/skills/triton/kernel-verifier/scripts/verify.py b/skills/triton/kernel-verifier/scripts/verify.py index cd7c47bf..1ba83e8b 100644 --- a/skills/triton/kernel-verifier/scripts/verify.py +++ b/skills/triton/kernel-verifier/scripts/verify.py @@ -12,7 +12,6 @@ import json import os import sys -import subprocess import traceback @@ -20,7 +19,7 @@ # 精度判定常量(dtype 无关) MAX_ERROR_CAP = 0.1 -REQUIRED_MATCHED_RATIO = 0.98 +REQUIRED_MATCHED_RATIO = 0.9 class AccuracyError(AssertionError): @@ -478,7 +477,7 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ "--verify_dir", default=".", help="验证目录,包含 {op_name}_torch.py 和 {op_name}_triton_ascend_impl.py(默认当前目录)", ) - parser.add_argument("--timeout", type=int, default=900, help="超时秒数(默认 900)") + parser.add_argument("--timeout", type=int, default=900, help="超时秒数(默认 900,已忽略:当前为同进程模式)") parser.add_argument( "--triton_impl_name", default="triton_ascend_impl", help="Triton 实现模块名(不含 op_name 前缀,默认 triton_ascend_impl)", @@ -487,10 +486,6 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ "--output", default=None, help="验证结果 JSON 输出路径(默认 {verify_dir}/verify_result.json)", ) - parser.add_argument( - "--_run", action="store_true", - help=argparse.SUPPRESS, # 内部参数:子进程模式,直接执行验证 - ) args = parser.parse_args() verify_dir = os.path.abspath(args.verify_dir) @@ -498,45 +493,12 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ print(f"错误: 验证目录不存在: {verify_dir}", file=sys.stderr) sys.exit(1) - if args._run: - # 子进程模式:直接执行验证逻辑 - try: - passed, total = verify_implementations( - args.op_name, verify_dir, args.triton_impl_name, args.output - ) - except Exception as e: - print(f"{e}", file=sys.stderr) - traceback.print_exc() - sys.exit(1) - # 策略 A:passed < total → exit 1 - sys.exit(0 if passed == total and total > 0 else 1) - else: - # 主进程模式:启动子进程执行验证,超时后 kill 整个进程树 - cmd = [ - sys.executable, os.path.abspath(__file__), - "--op_name", args.op_name, - "--verify_dir", verify_dir, - "--triton_impl_name", args.triton_impl_name, - "--_run", - ] - if args.output: - cmd.extend(["--output", args.output]) - try: - proc = subprocess.Popen( - cmd, - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - ) - stdout, stderr = proc.communicate(timeout=args.timeout) - - sys.stdout.buffer.write(stdout) - sys.stdout.buffer.flush() - sys.stderr.buffer.write(stderr) - sys.stderr.buffer.flush() - sys.exit(proc.returncode) - - except subprocess.TimeoutExpired: - proc.kill() - proc.wait() - print(f"验证超时({args.timeout}秒),已终止子进程", file=sys.stderr) - sys.exit(1) \ No newline at end of file + try: + passed, total = verify_implementations( + args.op_name, verify_dir, args.triton_impl_name, args.output + ) + except Exception as e: + print(f"{e}", file=sys.stderr) + traceback.print_exc() + sys.exit(1) + sys.exit(0 if passed == total and total > 0 else 1) \ No newline at end of file From a7bf0ba268f3865c2ea4aa7506a98e1ee53b41a0 Mon Sep 17 00:00:00 2001 From: w00934874 Date: Wed, 13 May 2026 14:58:31 +0800 Subject: [PATCH 6/9] =?UTF-8?q?[triton]=20verify=E8=AF=84=E6=B5=8B?= =?UTF-8?q?=E6=A0=87=E5=87=86=E6=9B=B4=E6=94=B9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- skills/triton/kernel-verifier/SKILL.md | 119 +++- .../triton/kernel-verifier/scripts/verify.py | 529 ++++++++++++++++-- 2 files changed, 606 insertions(+), 42 deletions(-) diff --git a/skills/triton/kernel-verifier/SKILL.md b/skills/triton/kernel-verifier/SKILL.md index 5b9e593b..1ef242f4 100644 --- a/skills/triton/kernel-verifier/SKILL.md +++ b/skills/triton/kernel-verifier/SKILL.md @@ -16,6 +16,61 @@ argument-hint: > 你是一个内核代码验证专家。你的任务是按照标准验证流程,创建验证项目并运行,检查生成的算子代码是否能正确编译运行且与参考实现的输出一致。验证通过后,执行性能测试并收集性能数据。 +## 验证分类与判定标准(概览) + +verify.py 按"`--non-compute` 开关 + 输入 dtype + 输出 dtype"分流到 **5 类判定路径**,详尽阈值表见文末"精度阈值说明"。 + +### 输入类型推断 + +从实际传入对象推断(KernelBench / NPUKernelBench 通用): + +1. 存在 `torch.Tensor` 输入 → 取所有 tensor 中**最高精度 dtype**(fp64 > fp32 > fp16 > bf16 > fp8 > int64 > int32 > int16 > int8/uint8 > bool) +2. 否则存在 `list/tuple of Tensor`(tensor_list)→ 取首个 tensor_list 首元素 dtype +3. 否则视为**无 tensor 输入**(最严路径) + +输入 dtype 落到浮点(含 fp8/复数)→ `input_type=float`;落到整型(含 bool)→ `input_type=int`。 + +### 五类判定决策矩阵 + +| 类别 | 输入 type | 输出 dtype | `--non-compute` | 误差要求 | +|---|---|---|---|---| +| **非计算类** | 任意 | 任意 | **是** | 二进制完全一致(view-as-int 比对,含 NaN bit pattern) | +| **bool 输出)** | 任意 | bool | 否 | `torch.equal` 严格相等 | +| **整数计算类** | int / no_tensor | int | 否 | `\|actual − golden\| == 0` | +| **量化计算类 fp→int** | float | int | 否 | `\|actual − golden\| <= 1` | +| **浮点计算类** | 任意 | float | 否 | 三项 AND(见下) | + +### 浮点计算类:三项整体判定(AND) + +1. **max_error_cap**:所有 finite 元素满足 `|diff| <= atol + rtol·|golden|`(dtype-aware,100% 通过) +2. **matched_ratio ≥ 0.9**:元素级匹配,小值域 `|golden| fp32 > fp16 > bf16 > fp8 > int64 > int32 > int16 > int8 > bool + """ + import torch + rank = { + torch.float64: 100, + torch.float32: 90, + torch.float16: 80, + torch.bfloat16: 70, + torch.int64: 50, + torch.int32: 40, + torch.int16: 30, + torch.int8: 20, + torch.uint8: 20, + torch.bool: 10, + } + for name in ("float8_e4m3fn", "float8_e4m3", "float8_e5m2fn", "float8_e5m2"): + dt = getattr(torch, name, None) + if dt is not None: + rank[dt] = 60 + return rank + + +_DTYPE_RANK = None + + +def _dtype_rank(dtype): + global _DTYPE_RANK + if _DTYPE_RANK is None: + _DTYPE_RANK = _build_dtype_rank() + return _DTYPE_RANK.get(dtype, 0) + + +def _is_int_like_dtype(dtype): + """判断 dtype 属于"整型类"输入(含 bool;不含浮点/复数)。""" + import torch + if dtype is None: + return False + if dtype == torch.bool: + return True + return (not dtype.is_floating_point) and (not dtype.is_complex) + + +def _infer_input_type(inputs): + """从 inputs 推断输入类型,返回 ("float" | "int" | "no_tensor", input_dtype | None)。 + + 判定优先级(KernelBench / NPUKernelBench 统一处理): + 1. 若存在 torch.Tensor 输入:取所有 tensor 中最高精度 dtype 作为输入类型 + 2. 若不存在 tensor,但存在 list/tuple of Tensor(tensor_list):取第一个 tensor_list 的首元素 dtype + 3. 其他情况(全为标量 attr / 无输入):返回 ("no_tensor", None) + + bool 输入归到 "int" 类(按规则:bool 输出单独处理;bool 输入与 int 同等对待)。 + """ + import torch + tensors = [x for x in inputs if isinstance(x, torch.Tensor)] + source = None + candidate_dtypes = [] + if tensors: + candidate_dtypes = [t.dtype for t in tensors] + top_dtype = max(candidate_dtypes, key=_dtype_rank) + source = "tensor" + else: + tensor_lists = [ + x for x in inputs + if isinstance(x, (list, tuple)) and len(x) > 0 + and all(isinstance(e, torch.Tensor) for e in x) + ] + if tensor_lists: + top_dtype = tensor_lists[0][0].dtype + candidate_dtypes = [top_dtype] + source = "tensor_list" + else: + print( + " [输入类型判定] 来源=无 tensor 输入(全 attr 或空)," + "input_type=no_tensor", + file=sys.stderr, + ) + return "no_tensor", None + + input_type = "int" if _is_int_like_dtype(top_dtype) else "float" + print( + f" [输入类型判定] 来源={source},候选 dtypes={[str(dt) for dt in candidate_dtypes]}," + f"最高精度={top_dtype},input_type={input_type}", + file=sys.stderr, + ) + return input_type, top_dtype + + def resolve_input_provider(torch_module): """解析任务文件的输入提供方式。""" if hasattr(torch_module, "get_input_groups"): @@ -139,8 +263,128 @@ def resolve_input_provider(torch_module): ) -def compare(fw_out, impl_out, data_type): - """对比框架输出和实现输出""" +def _compare_binary_exact(fw_out, impl_out, data_type): + """非计算类:二进制完全一致比对。 + + - 浮点 dtype:通过 view-as-int 比较底层 bit pattern,可识别 NaN payload 差异 + - 整型 / bool:直接 torch.equal + - 复数:实部/虚部分别 view-as-int 比较 + """ + import torch + + fw = fw_out.contiguous().detach().cpu() + impl = impl_out.contiguous() + if isinstance(impl, torch.Tensor): + impl = impl.detach().cpu() + else: + raise AssertionError(f"非计算类实现输出必须是 Tensor,实际为 {type(impl).__name__}") + + if fw.shape != impl.shape: + raise AssertionError( + f"非计算类验证失败,输出形状不一致: framework={fw.shape}, impl={impl.shape}" + ) + if fw.dtype != impl.dtype: + raise AssertionError( + f"非计算类验证失败,输出 dtype 不一致: framework={fw.dtype}, impl={impl.dtype}" + ) + + def _view_int_dtype(dt): + if dt in (torch.float64, torch.complex64): + return torch.int64 + if dt in (torch.float32,): + return torch.int32 + if dt in (torch.float16, torch.bfloat16): + return torch.int16 + for name in ("float8_e4m3fn", "float8_e4m3", "float8_e5m2fn", "float8_e5m2"): + fp8 = getattr(torch, name, None) + if fp8 is not None and dt == fp8: + return torch.int8 + return None + + if fw.dtype.is_complex: + fw_real_bits = torch.view_as_real(fw) + impl_real_bits = torch.view_as_real(impl) + view_dt = _view_int_dtype(torch.float32) if fw.dtype == torch.complex64 else torch.int64 + equal = torch.equal(fw_real_bits.view(view_dt), impl_real_bits.view(view_dt)) + elif fw.dtype.is_floating_point: + view_dt = _view_int_dtype(fw.dtype) + if view_dt is None: + raise AssertionError(f"非计算类不支持的浮点 dtype: {fw.dtype}") + equal = torch.equal(fw.view(view_dt), impl.view(view_dt)) + else: + equal = torch.equal(fw, impl) + + if equal: + return + + if fw.dtype.is_floating_point and not fw.dtype.is_complex: + view_dt = _view_int_dtype(fw.dtype) + fw_bits = fw.view(view_dt).flatten() + impl_bits = impl.view(view_dt).flatten() + diff_mask = fw_bits != impl_bits + violation_count = int(diff_mask.sum().item()) + violation_idx = torch.where(diff_mask)[0] + num_to_show = min(10, len(violation_idx)) + detail = f"前 {num_to_show} 个 bit 不一致位置:\n" + fw_flat = fw.flatten() + impl_flat = impl.flatten() + for i in range(num_to_show): + idx = violation_idx[i].item() + detail += ( + f" 位置[{idx}]: framework={fw_flat[idx].item()} " + f"(bits=0x{fw_bits[idx].item() & ((1 << view_dt.itemsize * 8) - 1):x}), " + f"impl={impl_flat[idx].item()} " + f"(bits=0x{impl_bits[idx].item() & ((1 << view_dt.itemsize * 8) - 1):x})\n" + ) + else: + fw_flat = fw.flatten() + impl_flat = impl.flatten() + diff_mask = fw_flat != impl_flat + violation_count = int(diff_mask.sum().item()) + violation_idx = torch.where(diff_mask)[0] + num_to_show = min(10, len(violation_idx)) + detail = f"前 {num_to_show} 个不一致位置:\n" + for i in range(num_to_show): + idx = violation_idx[i].item() + detail += ( + f" 位置[{idx}]: framework={fw_flat[idx].item()}, " + f"impl={impl_flat[idx].item()}\n" + ) + + metrics = { + "category": "non_compute", + "dtype": str(data_type), + "violation_count": violation_count, + "total_elements": int(fw.numel()), + } + raise AccuracyError( + f"验证失败 dtype={data_type} (非计算类,要求二进制完全一致): " + f"{violation_count}/{fw.numel()} 元素不一致\n{detail}", + metrics, + ) + + +def compare(fw_out, impl_out, data_type, input_type=None, input_dtype=None, non_compute=False): + """对比框架输出和实现输出。 + + Args: + fw_out: 框架(金标准)输出 Tensor + impl_out: 被测实现输出 Tensor + data_type: 输出 dtype(与 fw_out.dtype 一致) + input_type: 输入类型 "float" / "int" / "no_tensor" / None + 由 _infer_input_type() 推断得出,参与"输出整型时"的分流。 + input_dtype: 输入最高精度 dtype(由 _infer_input_type() 返回),仅用于诊断打印。 + non_compute: 若 True,强制走二进制完全一致路径(搬移 / Cast 等算子) + + 决策矩阵(non_compute=False 时): + | 输出 dtype | 输入类型 | 类别 | 判定 | + |-----------|------------------|---------------|--------------------| + | bool | 任意 | bool 输出 | torch.equal | + | int | int | 整数计算类 | |diff| == 0 | + | int | float | 量化类 | |diff| <= 1 | + | int | no_tensor | 整数计算类 | |diff| == 0(最严) | + | float | 任意 | 浮点计算类 | 三项判定(按输出 dtype)| + """ import torch fw_flat = fw_out.flatten().detach().cpu() impl_flat = impl_out.flatten() @@ -156,6 +400,17 @@ def compare(fw_out, impl_out, data_type): f"验证失败,输出形状不一致: framework={fw_flat.shape}, impl={impl_flat.shape}" ) + # 非计算类:二进制完全一致(先于其他判定,跳过 NaN/Inf/finite 过滤) + if non_compute: + print( + f" [评测模式] 模式=non_compute(非计算类)," + f"输入 dtype={input_dtype}({input_type}),输出 dtype={data_type};" + f"误差要求=二进制完全一致(view-as-int bit pattern 全等,含 NaN payload)", + file=sys.stderr, + ) + _compare_binary_exact(fw_out, impl_out, data_type) + return + fw_nan_mask = torch.isnan(fw_flat) impl_nan_mask = torch.isnan(impl_flat) if not torch.equal(fw_nan_mask, impl_nan_mask): @@ -191,35 +446,156 @@ def compare(fw_out, impl_out, data_type): fw_finite = fw_flat[finite_mask] impl_finite = impl_flat[finite_mask] + # bool 输出独立处理:严格相等 if fw_finite.dtype == torch.bool: + print( + f" [评测模式] 模式=bool_output(bool 输出)," + f"输入 dtype={input_dtype}({input_type}),输出 dtype={data_type};" + f"误差要求=torch.equal 严格相等(finite={finite_count}/{size})", + file=sys.stderr, + ) if not torch.equal(fw_finite, impl_finite): - raise AssertionError(f"验证失败,布尔值不匹配: dtype={data_type}") + diff_idx = torch.where(fw_finite != impl_finite)[0] + violation_count = int(diff_idx.numel()) + num_to_show = min(10, violation_count) + detail = f"前 {num_to_show} 个不一致位置:\n" + for i in range(num_to_show): + idx = diff_idx[i].item() + detail += ( + f" 位置[{idx}]: framework={fw_finite[idx].item()}, " + f"impl={impl_finite[idx].item()}\n" + ) + metrics = { + "category": "bool_output", + "dtype": str(data_type), + "violation_count": violation_count, + "total_finite": int(fw_finite.numel()), + } + raise AccuracyError( + f"验证失败 dtype={data_type} (bool 输出,要求严格相等): " + f"{violation_count}/{fw_finite.numel()} 元素不一致\n{detail}", + metrics, + ) return + # 输出整型:按 input_type 分流 + if _is_integer_dtype(fw_finite.dtype): + # input_type == "float" → 量化类 (|diff|<=1) + # input_type == "int" 或 "no_tensor" 或 None → 整数计算类 (|diff|==0,最严) + if input_type == "float": + print( + f" [评测模式] 模式=quant_fp_to_int(量化类 fp→int)," + f"输入 dtype={input_dtype}({input_type}),输出 dtype={data_type};" + f"误差要求=|actual - golden| <= 1(finite={finite_count}/{size})", + file=sys.stderr, + ) + diff = (fw_finite.to(torch.int64) - impl_finite.to(torch.int64)).abs() + violation_count = int((diff > 1).sum().item()) + if violation_count > 0: + max_diff = int(diff.max().item()) + violation_idx = torch.where(diff > 1)[0] + num_to_show = min(10, len(violation_idx)) + detail = f"前 {num_to_show} 个量化误差超限位置:\n" + for i in range(num_to_show): + idx = violation_idx[i].item() + detail += ( + f" 位置[{idx}]: framework={fw_finite[idx].item()}, " + f"impl={impl_finite[idx].item()}, " + f"|diff|={diff[idx].item()} (允许<=1)\n" + ) + metrics = { + "category": "quant_fp_to_int", + "dtype": str(data_type), + "input_type": input_type, + "max_abs_diff": max_diff, + "violation_count": violation_count, + "total_finite": int(diff.numel()), + "tolerance": 1, + } + raise AccuracyError( + f"验证失败 dtype={data_type} (量化类 fp->int,要求|diff|<=1): " + f"{violation_count}/{diff.numel()} 元素超限,max_abs_diff={max_diff}\n" + f"{detail}", + metrics, + ) + return + else: + # 整数计算类:严格相等 + print( + f" [评测模式] 模式=integer_compute(整数计算类)," + f"输入 dtype={input_dtype}({input_type}),输出 dtype={data_type};" + f"误差要求=|actual - golden| == 0(严格相等,finite={finite_count}/{size})", + file=sys.stderr, + ) + if not torch.equal(fw_finite, impl_finite): + diff = (fw_finite.to(torch.int64) - impl_finite.to(torch.int64)).abs() + violation_count = int((diff > 0).sum().item()) + max_diff = int(diff.max().item()) + violation_idx = torch.where(diff > 0)[0] + num_to_show = min(10, len(violation_idx)) + detail = f"前 {num_to_show} 个不一致位置:\n" + for i in range(num_to_show): + idx = violation_idx[i].item() + detail += ( + f" 位置[{idx}]: framework={fw_finite[idx].item()}, " + f"impl={impl_finite[idx].item()}, " + f"|diff|={diff[idx].item()}\n" + ) + metrics = { + "category": "integer_compute", + "dtype": str(data_type), + "input_type": input_type, + "max_abs_diff": max_diff, + "violation_count": violation_count, + "total_finite": int(diff.numel()), + "tolerance": 0, + } + raise AccuracyError( + f"验证失败 dtype={data_type} (整数计算类,要求严格相等): " + f"{violation_count}/{diff.numel()} 元素不一致,max_abs_diff={max_diff}\n" + f"{detail}", + metrics, + ) + return + if impl_finite.dtype != fw_finite.dtype: impl_finite = impl_finite.to(fw_finite.dtype) - # 执行 NPU Benchmark 精度验证 + # 输出浮点:按浮点精度标准执行(dtype-aware 三项判定) + sv_thr_pre, sv_err_pre, rel_thr_pre = get_limits(data_type) + atol_pre, rtol_pre = get_allclose_tols(data_type) + print( + f" [评测模式] 模式=float_compute(浮点计算类)," + f"输入 dtype={input_dtype}({input_type}),输出 dtype={data_type};" + f"误差要求=三项 AND:" + f"(1)max_error_cap |diff|<=atol+rtol*|golden| " + f"[atol={atol_pre:.3e}, rtol={rtol_pre:.3e}]," + f"(2)matched_ratio>={REQUIRED_MATCHED_RATIO} " + f"[小值域 sv_thr={sv_thr_pre:.3e}/sv_err={sv_err_pre:.3e}," + f"正常域 rel_thr={rel_thr_pre:.3e}]," + f"(3)MERE<{rel_thr_pre:.3e}(finite={finite_count}/{size})", + file=sys.stderr, + ) _check_accuracy_npu_benchmark(fw_finite, impl_finite, data_type) def _check_accuracy_npu_benchmark(golden, actual, data_type): - """执行 NPU Benchmark 精度验证(分类 + 三项判定)。 + """执行 NPU Benchmark 精度验证(三项判定)。 - 元素级 matched 定义: + 元素级 matched 定义(用于 #2 matched_ratio): - |golden| < small_value_threshold(小值域):|diff| <= small_value_error - 否则(正常值域):|diff| / (|golden| + 1e-7) <= rel_threshold 通过条件(三项 AND): - 1. max(|diff|) <= MAX_ERROR_CAP(0.1,dtype 无关的绝对误差上限) - 2. matched_ratio = sum(matched) / total_finite >= REQUIRED_MATCHED_RATIO(0.98) + 1. allclose: 所有元素满足 |diff| <= atol + rtol * |golden|(dtype-aware) + 2. matched_ratio = sum(matched) / total_finite >= REQUIRED_MATCHED_RATIO(0.9) 3. MERE < rel_threshold(对所有 finite 元素计算相对误差再取均值, 分母统一用 |golden| + 1e-7 防除零) Args: golden: 参考输出(金标准) actual: 被测实现输出 - data_type: 数据类型,用于获取对应的阈值三元组 + data_type: 数据类型,用于获取对应阈值 Raises: AccuracyError: 当精度验证未通过时,异常的 metrics 属性携带结构化指标 @@ -231,15 +607,16 @@ def _check_accuracy_npu_benchmark(golden, actual, data_type): actual_f = actual.float() sv_thr, sv_err, rel_thr = get_limits(data_type) + atol, rtol = get_allclose_tols(data_type) abs_diff = (actual_f - golden_f).abs() abs_golden = golden_f.abs() - # 分桶 + # 分桶(用于 #2 matched_ratio) small_mask = abs_golden < sv_thr normal_mask = ~small_mask - # 元素级 matched + # 元素级 matched(#2 口径) small_ok = abs_diff <= sv_err rel_err = abs_diff / (abs_golden + 1e-7) normal_ok = rel_err <= rel_thr @@ -250,6 +627,12 @@ def _check_accuracy_npu_benchmark(golden, actual, data_type): matched_ratio = matched_count / total_finite if total_finite > 0 else 1.0 max_abs_diff = abs_diff.max().item() if total_finite > 0 else 0.0 + # #1 allclose:逐元素判定,要求 100% 通过 + allclose_bound = atol + rtol * abs_golden + allclose_mask = abs_diff <= allclose_bound + allclose_violation_count = int((~allclose_mask).sum().item()) if total_finite > 0 else 0 + allclose_ok = allclose_violation_count == 0 + # MERE:对所有 finite 元素计算相对误差再取均值(分母统一 |golden| + 1e-7 防除零) normal_count = int(normal_mask.sum().item()) if total_finite > 0: @@ -259,9 +642,8 @@ def _check_accuracy_npu_benchmark(golden, actual, data_type): MERE = None mere_ok = True - cap_ok = max_abs_diff <= MAX_ERROR_CAP ratio_ok = matched_ratio >= REQUIRED_MATCHED_RATIO - is_pass = cap_ok and ratio_ok and mere_ok + is_pass = allclose_ok and ratio_ok and mere_ok if is_pass: return @@ -273,34 +655,49 @@ def _check_accuracy_npu_benchmark(golden, actual, data_type): "rel_threshold": rel_thr, "small_value_threshold": sv_thr, "small_value_error": sv_err, - "max_error_cap": MAX_ERROR_CAP, + "atol": atol, + "rtol": rtol, + "max_error_cap_violation_count": allclose_violation_count, "required_matched_ratio": REQUIRED_MATCHED_RATIO, "total_finite": total_finite, "matched_count": matched_count, "small_count": int(small_mask.sum().item()), "normal_count": normal_count, "checks": { - "max_error_cap": cap_ok, + "max_error_cap": allclose_ok, "required_matched_ratio": ratio_ok, "MERE": mere_ok, }, } - # 失败摘要 + 前 N 个 unmatched 位置(按所属桶注明判定标准) - unmatched_mask = ~matched_mask - unmatched_indices = torch.where(unmatched_mask)[0] - num_to_show = min(10, len(unmatched_indices)) - mere_str = f"{MERE:.6e}" if MERE is not None else "n/a" error_msg = ( f"验证失败 dtype={data_type}: " - f"max_abs_diff={max_abs_diff:.6e} (cap={MAX_ERROR_CAP}, ok={cap_ok}), " + f"max_error_cap_violations={allclose_violation_count}/{total_finite} " + f"(atol={atol:.6e}, rtol={rtol:.6e}, max_abs_diff={max_abs_diff:.6e}, ok={allclose_ok}), " f"matched_ratio={matched_ratio:.6f} (req>={REQUIRED_MATCHED_RATIO}, ok={ratio_ok}), " f"MERE={mere_str} (rel_thr={rel_thr:.6e}, ok={mere_ok}); " f"small_count={metrics['small_count']}, normal_count={normal_count}\n" ) - if num_to_show > 0: - error_msg += f"前 {num_to_show} 个未通过的位置:\n" + + # 仅在对应检查失败时打印各自的违例位置(前 N 个) + if not allclose_ok: + allclose_violation_indices = torch.where(~allclose_mask)[0] + num_to_show = min(10, len(allclose_violation_indices)) + error_msg += f"前 {num_to_show} 个 max_error_cap 违例位置:\n" + for i in range(num_to_show): + idx = allclose_violation_indices[i].item() + error_msg += ( + f" 位置[{idx}]: framework={golden[idx]:.6e}, " + f"impl={actual[idx]:.6e}, |diff|={abs_diff[idx]:.6e} " + f"(允许<=atol+rtol*|golden|={allclose_bound[idx]:.6e})\n" + ) + + if not ratio_ok: + unmatched_mask = ~matched_mask + unmatched_indices = torch.where(unmatched_mask)[0] + num_to_show = min(10, len(unmatched_indices)) + error_msg += f"前 {num_to_show} 个 matched 未通过位置:\n" for i in range(num_to_show): idx = unmatched_indices[i].item() if small_mask[idx].item(): @@ -324,13 +721,17 @@ def run_single_case( inputs, device, case_idx, - total_cases + total_cases, + non_compute=False, ): """验证单组输入。失败时抛出 AssertionError。""" import torch print(f" 测试第 {case_idx}/{total_cases} 组输入...", file=sys.stderr) + # 推断输入类型("float" / "int" / "no_tensor")→ 决定输出整型时走整数计算 vs 量化 + input_type, input_dtype = _infer_input_type(inputs) + inputs_for_impl = [ x.to(device) if isinstance(x, torch.Tensor) else x for x in inputs @@ -355,6 +756,11 @@ def run_single_case( f"framework={len(framework_output)}, impl={len(impl_output)}" ) + print( + f" [输出概览] 共 {len(framework_output)} 个输出,non_compute={non_compute}", + file=sys.stderr, + ) + for i, (fw_out, impl_out) in enumerate(zip(framework_output, impl_output)): if fw_out is None or impl_out is None: raise AssertionError( @@ -364,7 +770,15 @@ def run_single_case( if isinstance(fw_out, torch.Tensor) and isinstance(impl_out, torch.Tensor): try: data_type = fw_out.dtype - compare(fw_out, impl_out, data_type) + print( + f" [输出 {i}] shape={list(fw_out.shape)}, dtype={data_type}", + file=sys.stderr, + ) + compare( + fw_out, impl_out, data_type, + input_type=input_type, input_dtype=input_dtype, + non_compute=non_compute, + ) except AccuracyError as e: raise AccuracyError( f"[用例 {case_idx}/{total_cases}] {str(e)}", e.metrics @@ -373,11 +787,14 @@ def run_single_case( raise AssertionError(f"[用例 {case_idx}/{total_cases}] {str(e)}") from e -def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_impl", output_path=None): +def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_impl", output_path=None, non_compute=False): """验证框架实现和生成实现的结果一致性。 每个 shape 独立 try/except,全部跑完后写 verify_result.json。 + Args: + non_compute: 若 True,所有 case 走"非计算类"二进制完全一致判定(搬移/Cast 等算子) + Returns: (passed_cases, total_cases) """ @@ -419,7 +836,8 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ impl_model = ModelNew(*init_params).to(device) run_single_case( - framework_model, impl_model, inputs, device, case_idx, total_cases + framework_model, impl_model, inputs, device, case_idx, total_cases, + non_compute=non_compute, ) passed_cases += 1 except Exception as e: @@ -486,6 +904,10 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ "--output", default=None, help="验证结果 JSON 输出路径(默认 {verify_dir}/verify_result.json)", ) + parser.add_argument( + "--_run", action="store_true", + help=argparse.SUPPRESS, # 内部参数:子进程模式,直接执行验证 + ) args = parser.parse_args() verify_dir = os.path.abspath(args.verify_dir) @@ -493,12 +915,45 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ print(f"错误: 验证目录不存在: {verify_dir}", file=sys.stderr) sys.exit(1) - try: - passed, total = verify_implementations( - args.op_name, verify_dir, args.triton_impl_name, args.output - ) - except Exception as e: - print(f"{e}", file=sys.stderr) - traceback.print_exc() - sys.exit(1) - sys.exit(0 if passed == total and total > 0 else 1) \ No newline at end of file + if args._run: + # 子进程模式:直接执行验证逻辑 + try: + passed, total = verify_implementations( + args.op_name, verify_dir, args.triton_impl_name, args.output + ) + except Exception as e: + print(f"{e}", file=sys.stderr) + traceback.print_exc() + sys.exit(1) + # 策略 A:passed < total → exit 1 + sys.exit(0 if passed == total and total > 0 else 1) + else: + # 主进程模式:启动子进程执行验证,超时后 kill 整个进程树 + cmd = [ + sys.executable, os.path.abspath(__file__), + "--op_name", args.op_name, + "--verify_dir", verify_dir, + "--triton_impl_name", args.triton_impl_name, + "--_run", + ] + if args.output: + cmd.extend(["--output", args.output]) + try: + proc = subprocess.Popen( + cmd, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + ) + stdout, stderr = proc.communicate(timeout=args.timeout) + + sys.stdout.buffer.write(stdout) + sys.stdout.buffer.flush() + sys.stderr.buffer.write(stderr) + sys.stderr.buffer.flush() + sys.exit(proc.returncode) + + except subprocess.TimeoutExpired: + proc.kill() + proc.wait() + print(f"验证超时({args.timeout}秒),已终止子进程", file=sys.stderr) + sys.exit(1) \ No newline at end of file From 73558e2c055b7eef9f40463d1bc9d54f690a1d99 Mon Sep 17 00:00:00 2001 From: w00934874 Date: Fri, 15 May 2026 10:54:19 +0800 Subject: [PATCH 7/9] =?UTF-8?q?[triton]verify=E4=BF=AE=E5=A4=8D=E9=97=AE?= =?UTF-8?q?=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../triton/kernel-verifier/scripts/verify.py | 71 ++++++------------- 1 file changed, 23 insertions(+), 48 deletions(-) diff --git a/skills/triton/kernel-verifier/scripts/verify.py b/skills/triton/kernel-verifier/scripts/verify.py index 69896c6c..beba5d1f 100644 --- a/skills/triton/kernel-verifier/scripts/verify.py +++ b/skills/triton/kernel-verifier/scripts/verify.py @@ -17,9 +17,15 @@ ERROR_MSG_LIMIT = 2000 -# 精度判定常量(dtype 无关) -MAX_ERROR_CAP = 0.1 -REQUIRED_MATCHED_RATIO = 0.98 +REQUIRED_MATCHED_RATIO = 0.9 + +# allclose 判定阈值 (atol, rtol):|actual - golden| <= atol + rtol * |golden| +ALLCLOSE_TOLS_STR = { + "float32": (1e-3, 2**(-13)), # 2**(-13)=1.220703125e-4 + "float16": (5e-3, 2**(-10)), # 2**(-10)=9.765625e-4 + "bfloat16": (1e-2, 2**(-7)), # 2**(-7)=7.8125e-3 +} +ALLCLOSE_DEFAULT_TOLS = ALLCLOSE_TOLS_STR["float32"] class AccuracyError(AssertionError): @@ -895,7 +901,7 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ "--verify_dir", default=".", help="验证目录,包含 {op_name}_torch.py 和 {op_name}_triton_ascend_impl.py(默认当前目录)", ) - parser.add_argument("--timeout", type=int, default=900, help="超时秒数(默认 900,已忽略:当前为同进程模式)") + parser.add_argument("--timeout", type=int, default=900, help="超时秒数(已忽略:当前为同进程串行模式)") parser.add_argument( "--triton_impl_name", default="triton_ascend_impl", help="Triton 实现模块名(不含 op_name 前缀,默认 triton_ascend_impl)", @@ -905,8 +911,8 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ help="验证结果 JSON 输出路径(默认 {verify_dir}/verify_result.json)", ) parser.add_argument( - "--_run", action="store_true", - help=argparse.SUPPRESS, # 内部参数:子进程模式,直接执行验证 + "--non-compute", action="store_true", + help="非计算类算子(搬移 / Cast 等),所有 case 走二进制完全一致判定", ) args = parser.parse_args() @@ -915,45 +921,14 @@ def verify_implementations(op_name, verify_dir, triton_impl_name="triton_ascend_ print(f"错误: 验证目录不存在: {verify_dir}", file=sys.stderr) sys.exit(1) - if args._run: - # 子进程模式:直接执行验证逻辑 - try: - passed, total = verify_implementations( - args.op_name, verify_dir, args.triton_impl_name, args.output - ) - except Exception as e: - print(f"{e}", file=sys.stderr) - traceback.print_exc() - sys.exit(1) - # 策略 A:passed < total → exit 1 - sys.exit(0 if passed == total and total > 0 else 1) - else: - # 主进程模式:启动子进程执行验证,超时后 kill 整个进程树 - cmd = [ - sys.executable, os.path.abspath(__file__), - "--op_name", args.op_name, - "--verify_dir", verify_dir, - "--triton_impl_name", args.triton_impl_name, - "--_run", - ] - if args.output: - cmd.extend(["--output", args.output]) - try: - proc = subprocess.Popen( - cmd, - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - ) - stdout, stderr = proc.communicate(timeout=args.timeout) - - sys.stdout.buffer.write(stdout) - sys.stdout.buffer.flush() - sys.stderr.buffer.write(stderr) - sys.stderr.buffer.flush() - sys.exit(proc.returncode) - - except subprocess.TimeoutExpired: - proc.kill() - proc.wait() - print(f"验证超时({args.timeout}秒),已终止子进程", file=sys.stderr) - sys.exit(1) \ No newline at end of file + try: + passed, total = verify_implementations( + args.op_name, verify_dir, args.triton_impl_name, args.output, + non_compute=args.non_compute, + ) + except Exception as e: + print(f"{e}", file=sys.stderr) + traceback.print_exc() + sys.exit(1) + # 策略 A:passed < total → exit 1 + sys.exit(0 if passed == total and total > 0 else 1) \ No newline at end of file From bde88cf08f39d6986d6c4e83a8d77fa94a7896f4 Mon Sep 17 00:00:00 2001 From: w00934874 Date: Fri, 15 May 2026 11:57:08 +0800 Subject: [PATCH 8/9] =?UTF-8?q?=E4=BF=AE=E6=AD=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- skills/triton/kernel-verifier/scripts/verify.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/skills/triton/kernel-verifier/scripts/verify.py b/skills/triton/kernel-verifier/scripts/verify.py index beba5d1f..f83dfee6 100644 --- a/skills/triton/kernel-verifier/scripts/verify.py +++ b/skills/triton/kernel-verifier/scripts/verify.py @@ -22,8 +22,8 @@ # allclose 判定阈值 (atol, rtol):|actual - golden| <= atol + rtol * |golden| ALLCLOSE_TOLS_STR = { "float32": (1e-3, 2**(-13)), # 2**(-13)=1.220703125e-4 - "float16": (5e-3, 2**(-10)), # 2**(-10)=9.765625e-4 - "bfloat16": (1e-2, 2**(-7)), # 2**(-7)=7.8125e-3 + "float16": (9e-2, 2**(-10)), # 2**(-10)=9.765625e-4 + "bfloat16": (1e-1, 2**(-7)), # 2**(-7)=7.8125e-3 } ALLCLOSE_DEFAULT_TOLS = ALLCLOSE_TOLS_STR["float32"] From 6f884dca6c1a523bd7c4ca7b6fe0436965b12958 Mon Sep 17 00:00:00 2001 From: w00934874 Date: Fri, 15 May 2026 16:26:17 +0800 Subject: [PATCH 9/9] =?UTF-8?q?[triton]=20=E6=95=B4=E4=BD=93=E4=BF=AE?= =?UTF-8?q?=E6=94=B9verify=20skill=E6=8F=8F=E8=BF=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- agents/triton-ascend-coder.md | 2 +- skills/triton/kernel-verifier/SKILL.md | 277 +++++++++++-------------- 2 files changed, 117 insertions(+), 162 deletions(-) diff --git a/agents/triton-ascend-coder.md b/agents/triton-ascend-coder.md index 8fa12a61..a8189607 100644 --- a/agents/triton-ascend-coder.md +++ b/agents/triton-ascend-coder.md @@ -152,7 +152,7 @@ Agent 自身维护迭代状态,编排 "生成 → 验证 → Conductor 分析" ``` iteration = 0 -max_iterations = 5 +max_iterations = 20 history_attempts = [] previous_code = "" verifier_error = "" diff --git a/skills/triton/kernel-verifier/SKILL.md b/skills/triton/kernel-verifier/SKILL.md index 1ef242f4..6a4a6241 100644 --- a/skills/triton/kernel-verifier/SKILL.md +++ b/skills/triton/kernel-verifier/SKILL.md @@ -16,54 +16,101 @@ argument-hint: > 你是一个内核代码验证专家。你的任务是按照标准验证流程,创建验证项目并运行,检查生成的算子代码是否能正确编译运行且与参考实现的输出一致。验证通过后,执行性能测试并收集性能数据。 -## 验证分类与判定标准(概览) +## 精度判定规则 -verify.py 按"`--non-compute` 开关 + 输入 dtype + 输出 dtype"分流到 **5 类判定路径**,详尽阈值表见文末"精度阈值说明"。 +> 本节是 verify.py 精度判定的**唯一权威说明**。所有阈值、决策矩阵、前置检查只在此处定义;下游章节(Step 2/3)只引用、不重复。 +> (别名:精度阈值说明 / 验证分类与判定标准) -### 输入类型推断 +verify.py 按"`--non-compute` 开关 + 输入 dtype + 输出 dtype"分流到 **5 类判定路径**。 -从实际传入对象推断(KernelBench / NPUKernelBench 通用): +### 1. 输入类型推断(KernelBench / NPUKernelBench 统一) -1. 存在 `torch.Tensor` 输入 → 取所有 tensor 中**最高精度 dtype**(fp64 > fp32 > fp16 > bf16 > fp8 > int64 > int32 > int16 > int8/uint8 > bool) +从实际传入对象推断(不依赖 task 文件结构化 spec): + +1. 存在 `torch.Tensor` 输入 → 取所有 tensor 中**最高精度 dtype** + (例:输入为 `[fp16 tensor, fp32 tensor, int64 tensor]` → 取 fp32) 2. 否则存在 `list/tuple of Tensor`(tensor_list)→ 取首个 tensor_list 首元素 dtype -3. 否则视为**无 tensor 输入**(最严路径) +3. 否则视为**无 tensor 输入**(`no_tensor`,最严路径) + +**dtype 优先级**: + +精度从高到低排列:float64 > float32 > float16 > bfloat16 > float8_e4m3 / float8_e5m2 > int64 > int32 > int16 > int8 / uint8 > bool + +**`input_type` 二分类**(用于 §2 决策矩阵分流): + +- `input_type = float`:输入最高精度 dtype 属于浮点族(float64/32/16、bfloat16、float8_e4m3/e5m2、复数 complex64/128) +- `input_type = int`:输入最高精度 dtype 属于整型族(int64/32/16/8、uint8、bool) +- `input_type = no_tensor`:无任何 tensor 输入(走最严判定) -输入 dtype 落到浮点(含 fp8/复数)→ `input_type=float`;落到整型(含 bool)→ `input_type=int`。 +> **关于 bool 的两处特殊性**(避免混淆): +> - bool 作**输入**:归入 `int`,与 int 系输入等同;分流时只看输出 dtype(如输出 fp 走浮点判定) +> - bool 作**输出**:不进入 input_type 分流,直接走 §2 "bool 输出"路径(`torch.equal` 严格相等) -### 五类判定决策矩阵 +### 2. 五类判定决策矩阵 | 类别 | 输入 type | 输出 dtype | `--non-compute` | 误差要求 | |---|---|---|---|---| | **非计算类** | 任意 | 任意 | **是** | 二进制完全一致(view-as-int 比对,含 NaN bit pattern) | -| **bool 输出)** | 任意 | bool | 否 | `torch.equal` 严格相等 | +| **bool 输出** | 任意 | bool | 否 | `torch.equal` 严格相等 | | **整数计算类** | int / no_tensor | int | 否 | `\|actual − golden\| == 0` | | **量化计算类 fp→int** | float | int | 否 | `\|actual − golden\| <= 1` | -| **浮点计算类** | 任意 | float | 否 | 三项 AND(见下) | +| **浮点计算类** | 任意 | float | 否 | 三项 AND(见 §4) | + +### 3. 比对前置检查(按顺序,任一失败即判 fail) + +1. 形状必须一致 +2. NaN 位置必须完全一致(mask 按位相等) +3. Inf 位置和符号必须完全一致 +4. `bool` dtype:要求 `torch.equal` 完全相等,不进入精度判定 +5. 仅在 `finite_mask`(双方都 finite)上做精度计算;dtype 不一致时 impl 会被 cast 到 golden 的 dtype + +### 4. 浮点计算类:三项 AND 整体判定 -### 浮点计算类:三项整体判定(AND) +#### 4.1 元素级 matched 定义(分桶) -1. **max_error_cap**:所有 finite 元素满足 `|diff| <= atol + rtol·|golden|`(dtype-aware,100% 通过) -2. **matched_ratio ≥ 0.9**:元素级匹配,小值域 `|golden|= small_value_threshold`: + `matched[i] = (|actual[i] - golden[i]| / (|golden[i]| + 1e-7) <= rel_threshold)` -| dtype | atol | rtol | sv_thr | sv_err | rel_thr | -|---|---|---|---|---|---| -| float32 | 2e-4 | 2⁻¹³ ≈ 1.22e-4 | 2⁻¹⁴ ≈ 6.10e-5 | 2⁻³⁰ ≈ 9.31e-10 | 2⁻¹³ ≈ 1.22e-4 | -| float16 | 1e-2 | 2⁻¹⁰ ≈ 9.77e-4 | 2⁻¹¹ ≈ 4.88e-4 | 2⁻¹⁶ ≈ 1.53e-5 | 2⁻¹⁰ ≈ 9.77e-4 | -| bfloat16 | 1e-2 | 2⁻⁷ ≈ 7.81e-3 | 2⁻⁸ ≈ 3.91e-3 | 2⁻¹⁶ ≈ 1.53e-5 | 2⁻⁷ ≈ 7.81e-3 | +> 计算前两侧统一升 float32,避免低精度 dtype 自身误差污染。 +> 分母 `+1e-7` 仅为保险——正常域里 `|golden| >= sv_thr ≫ 1e-7`。 -### 比对前置检查(任一失败即 fail,先于上述判定) +#### 4.2 三项通过条件(AND,全部满足才算通过) -1. 形状一致 -2. NaN mask 完全一致 -3. Inf 位置 + 符号完全一致 -4. 仅在 `finite_mask`(双方都 finite)上做后续判定;dtype 不一致时 impl 会被 cast 到 golden dtype +1. **`max_error_cap`**:所有 finite 元素满足 `|diff| <= atol + rtol * |golden|`(dtype-aware,要求 100% 通过) +2. **`required_matched_ratio`**:`sum(matched) / total_finite >= 0.9` +3. **`MERE`**:对所有 finite 元素计算 `rel_err = |diff| / (|golden| + 1e-7)` 再取均值,要求 `MERE < rel_threshold`。当 `total_finite == 0` 时本项自动通过。 -### 运行时诊断 +#### 4.3 阈值表 + +**matched_mask 与 MERE 阈值**(沿用 NPU Benchmark 标准): + +| 数据类型 | small_value_threshold | small_value_error | rel_threshold (= MERE 上限) | +|---|---|---|---| +| `float16` | 2⁻¹¹ ≈ 4.88e-4 | 2⁻¹⁶ ≈ 1.53e-5 | 2⁻¹⁰ ≈ 9.77e-4 | +| `bfloat16` | 2⁻⁸ ≈ 3.91e-3 | 2⁻¹⁶ ≈ 1.53e-5 | 2⁻⁷ ≈ 7.81e-3 | +| `float32` | 2⁻¹⁴ ≈ 6.10e-5 | 2⁻³⁰ ≈ 9.31e-10 | 2⁻¹³ ≈ 1.22e-4 | +| `hifloat32` | 2⁻¹² ≈ 2.44e-4 | 2⁻²⁸ ≈ 3.73e-9 | 2⁻¹¹ ≈ 4.88e-4 | +| `float8_e4m3` | 2⁻⁴ = 0.0625 | 2⁻⁶ ≈ 0.015625 | 2⁻³ = 0.125 | +| `float8_e5m2` | 2⁻³ = 0.125 | 2⁻⁵ = 0.03125 | 2⁻² = 0.25 | +| 其他 dtype(fallback) | 2⁻¹⁴ | 2⁻³⁰ | 2⁻¹³ | + +**max_error_cap 阈值**(`|diff| <= atol + rtol * |golden|`): + +| 数据类型 | atol | rtol | +|---|---|---| +| `float16` | 9e-2 | 2⁻¹⁰ ≈ 9.77e-4 | +| `bfloat16` | 1e-1 | 2⁻⁷ ≈ 7.81e-3 | +| `float32` | 1e-3 | 2⁻¹³ ≈ 1.22e-4 | +| 其他 dtype(fallback) | 1e-3 | 2⁻¹³ | + +### 5. 运行时诊断输出 verify.py 在每个 case 会向 stderr 打印: + - `[输入类型判定] 来源=...,候选 dtypes=...,最高精度=...,input_type=...` - `[评测模式] 模式=...,输入 dtype=...,输出 dtype=...,误差要求=...` @@ -178,10 +225,11 @@ python3 /path/to/kernel-verifier/scripts/verify.py \ | `--op_name` | 是 | 算子名称,与文件名前缀对应 | | `--verify_dir` | 否 | 验证目录路径,默认当前目录 | | `--triton_impl_name` | 否 | Triton 实现模块名(不含 `{op_name}_` 前缀),默认 `triton_ascend_impl` | -| `--timeout` | 否 | 超时秒数,默认 900 | +| `--timeout` | 否 | 超时秒数,默认 900(⚠️ 当前已忽略:脚本为同进程串行模式,未实现超时强制中止) | +| `--output` | 否 | 验证结果 JSON 输出路径,默认 `{verify_dir}/verify_result.json` | | `--non-compute` | 否 | 适用于非计算类算子(不做数值运算、只对张量进行形状变换、维度重排、切分拼接、索引、类型转换等数据重组操作的算子,常见如 Reshape、Transpose、Concat、Split、Gather、Cast、Pad 等),强制走二进制完全一致判定 | -**超时设置**:默认 900 秒,复杂算子可适当增加。 +**超时设置**:`--timeout` 参数当前已忽略(同进程串行模式),保留仅为兼容旧调用方。 **注意事项**: - 禁止自己编写 Python 代码来测试算子(如手动 import 并 forward 比较) @@ -215,13 +263,49 @@ verify.py 会在 `verify_dir` 下生成 `verify_result.json`(或 `--output` } ``` +**精度失败时的 `metrics` 字段**:当 `error_type == "AccuracyError"`(浮点三项判定未通过)时,`failures[*]` 会带上结构化 `metrics`,便于下游分类失败原因(max_error_cap 违例 / 离群点过多 / 平均误差偏大): + +```json +{ + "case_idx": 1, + "input_desc": [...], + "error_type": "AccuracyError", + "error_msg": "...", + "metrics": { + "matched_ratio": 0.95, + "max_abs_diff": 0.2, + "MERE": 2.0e-4, + "rel_threshold": 1.22e-4, + "small_value_threshold": 6.10e-5, + "small_value_error": 9.31e-10, + "atol": 1.0e-3, + "rtol": 1.22e-4, + "max_error_cap_violation_count": 12, + "required_matched_ratio": 0.9, + "total_finite": 1000, + "matched_count": 950, + "small_count": 0, + "normal_count": 1000, + "checks": { + "max_error_cap": false, + "required_matched_ratio": false, + "MERE": false + } + } +} +``` + +`checks` 三个布尔位标记每项判定是否独立通过。阈值定义见上文 §精度判定规则 §4。 + +非浮点类失败(`non_compute` / `bool_output` / `integer_compute` / `quant_fp_to_int`)的 `metrics` 字段较简单,含 `category` / `violation_count` / `total_*` 等基本计数。 + **多 shape 行为**:每个 shape 独立 try/except,失败不中止后续 shape;全部跑完才落盘并退出。 **退出码语义(策略 A:严格)**: -- `passed_cases == total_cases` → exit 0,`verifier_result = true` -- `passed_cases < total_cases` → exit 1,`verifier_result = false`,`verifier_error` 应读取 `verify_result.json.failures` 的**全部条目**(不是第一个),汇总后提交给 Conductor。 +- `passed_cases == total_cases` 且 `total_cases > 0` → exit 0,`verifier_result = true` +- 否则(`passed_cases < total_cases`,或 `total_cases == 0`)→ exit 1,`verifier_result = false`,`verifier_error` 应读取 `verify_result.json.failures` 的**全部条目**(不是第一个),汇总后提交给 Conductor。 -**超时**:脚本输出 `"验证超时"` 且退出码为 1 → `verifier_error = "验证超时({timeout}秒)"`。 +**超时**:当前 `--timeout` 已被忽略,脚本不会主动中止;不会输出"验证超时"。 --- @@ -383,135 +467,6 @@ benchmark.py 启动时按 `--triton_impl_name` 推导对应的 verify_result 文 --- -## 精度阈值说明 - -验证按"输入类型 + 输出 dtype + --non-compute 开关"分流到四类判定路径。 - -### 输入类型判定(KernelBench / NPUKernelBench 统一推断) - -从实际传入的 Python 对象推断(不依赖 task 文件结构化 spec): - -1. 若 `inputs` 中存在 `torch.Tensor`:取所有 tensor 中**最高精度 dtype**作为输入类型 -2. 否则若存在 `list/tuple of Tensor`(tensor_list):取首个 tensor_list 的首元素 dtype -3. 否则视为**无 tensor 输入**,触发"最严格"路径 - -**dtype 优先级表**(值越大精度越高): - -| dtype | rank | -|-------|------| -| float64 | 100 | -| float32 | 90 | -| float16 | 80 | -| bfloat16 | 70 | -| float8_e4m3 / float8_e5m2 | 60 | -| int64 | 50 | -| int32 | 40 | -| int16 | 30 | -| int8 / uint8 | 20 | -| bool | 10 | - -输入 dtype 落到浮点(含 fp8/complex)→ 标记为 `float`;落到整型(含 bool)→ 标记为 `int`。 - -### 算子四分类(决策矩阵) - -| --non-compute | 输出 dtype | 输入类型 | 类别 | 判定 | -|---|---|---|---|---| -| 是 | 任意 | 任意 | **非计算类** | 二进制完全一致(view-as-int 比对,含 NaN bit pattern)| -| 否 | bool | 任意 | **bool 输出** | `torch.equal` 严格相等 | -| 否 | int | int | **整数计算类** | `|diff| == 0` | -| 否 | int | float | **量化计算类** | `|diff| <= 1` | -| 否 | int | no_tensor | 整数计算类(最严)| `|diff| == 0` | -| 否 | float | 任意 | **浮点计算类** | 三项判定(按输出 dtype 阈值)| - -### 浮点计算类:三项整体判定(NPU Benchmark 标准) - -### 元素级 matched 定义(分类) - -对每个 finite 元素 `i`,按 `|golden[i]|` 落入的类别分别判定: - -- **小值域** `|golden[i]| < small_value_threshold`: - `matched[i] = (|actual[i] - golden[i]| <= small_value_error)` -- **正常域** `|golden[i]| >= small_value_threshold`: - `matched[i] = (|actual[i] - golden[i]| / (|golden[i]| + 1e-7) <= rel_threshold)` - -> 计算前两侧统一升 float32,避免低精度 dtype 自身误差污染。 -> 分母 `+1e-7` 仅为保险——正常域里 `|golden| >= small_value_threshold ≫ 1e-7`。 - -### 通过条件(三项 AND,全部满足才算通过) - -1. **`max_error_cap`**:`max(|actual - golden|) <= 0.1`(全局绝对误差上限,dtype 无关) -2. **`required_matched_ratio`**:`sum(matched) / total_finite >= 0.9` -3. **`MERE`**:对**所有 finite 元素**计算 `rel_err = |diff| / (|golden| + 1e-7)` 再取均值,要求 `MERE < rel_threshold`。当 `total_finite == 0` 时本项自动通过。 - -### 阈值表 - -**matched_mask 与 MERE 阈值**(沿用 NPU Benchmark 标准): - -| 数据类型 | small_value_threshold | small_value_error | rel_threshold (= MERE 上限) | -|---|---|---|---| -| `float16` | 2⁻¹¹ ≈ 4.88e-4 | 2⁻¹⁶ ≈ 1.53e-5 | 2⁻¹⁰ ≈ 9.77e-4 | -| `bfloat16` | 2⁻⁸ ≈ 3.91e-3 | 2⁻¹⁶ ≈ 1.53e-5 | 2⁻⁷ ≈ 7.81e-3 | -| `float32` | 2⁻¹⁴ ≈ 6.10e-5 | 2⁻³⁰ ≈ 9.31e-10 | 2⁻¹³ ≈ 1.22e-4 | -| `hifloat32` | 2⁻¹² ≈ 2.44e-4 | 2⁻²⁸ ≈ 3.73e-9 | 2⁻¹¹ ≈ 4.88e-4 | -| `float8_e4m3` | 2⁻⁴ = 0.0625 | 2⁻⁶ ≈ 0.015625 | 2⁻³ = 0.125 | -| `float8_e5m2` | 2⁻³ = 0.125 | 2⁻⁵ = 0.03125 | 2⁻² = 0.25 | -| 其他 dtype(fallback) | 2⁻¹⁴ | 2⁻³⁰ | 2⁻¹³ | - -**max_error_cap 阈值**(`|diff| <= atol + rtol * |golden|`): - -| 数据类型 | atol | rtol | -|---|---|---| -| `float16` | 5e-3 | 2⁻¹⁰ ≈ 9.77e-4 | -| `bfloat16` | 1e-2 | 2⁻⁷ ≈ 7.81e-3 | -| `float32` | 2e-5 | 2⁻¹³ ≈ 1.22e-4 | -| 其他 dtype(fallback) | 2e-5 | 2⁻¹³ | - -### 失败时的 JSON 输出 - -`failures[*]` 在精度不达标时会带上结构化 `metrics`: - -```json -{ - "case_idx": 1, - "input_desc": [...], - "error_type": "AccuracyError", - "error_msg": "...", - "metrics": { - "matched_ratio": 0.95, - "max_abs_diff": 0.2, - "MERE": 2.0e-4, - "rel_threshold": 1.22e-4, - "small_value_threshold": 6.10e-5, - "small_value_error": 9.31e-10, - "atol": 2.0e-5, - "rtol": 1.22e-4, - "max_error_cap_violation_count": 12, - "required_matched_ratio": 0.9, - "total_finite": 1000, - "matched_count": 950, - "small_count": 0, - "normal_count": 1000, - "checks": { - "max_error_cap": false, - "required_matched_ratio": false, - "MERE": false - } - } -} -``` - -`checks` 三个布尔位标记每项判定是否独立通过,下游可直接据此分类失败原因(max_error_cap 违例 / 离群点过多 / 平均误差偏大)。 - -### 比对前置检查(按顺序,任一失败即判 fail) - -1. 形状必须一致 -2. NaN 位置必须完全一致(mask 按位相等) -3. Inf 位置和符号必须完全一致 -4. `bool` dtype:要求 `torch.equal` 完全相等,不进入精度判定 -5. 仅在 `finite_mask` 上做精度计算;当 dtype 不一致时 impl 会被 cast 到 golden 的 dtype - ---- - ## 脚本位置 验证脚本位于本 skill 的 `scripts/` 目录: @@ -524,5 +479,5 @@ benchmark.py 启动时按 `--triton_impl_name` 推导对应的 verify_result 文 **CLI 参数**: - `validate_triton_impl.py`: ``, `[--json]` -- `verify.py`: `--op_name`, `--verify_dir`, `--triton_impl_name`, `--timeout`, `--output` +- `verify.py`: `--op_name`, `--verify_dir`, `--triton_impl_name`, `--timeout`, `--output`, `--non-compute` - `benchmark.py`: `--op_name`, `--verify_dir`, `--triton_impl_name`, `--warmup`, `--repeats`, `--output`, `--skip_framework`, `--framework_latency_ms`, `--verify_not_required` \ No newline at end of file