【Bug已解决】AssertionError for FP8 Benchmarks 解决方案

发布时间:2026/8/1 15:25:34

【Bug已解决】AssertionError for FP8 Benchmarks 解决方案 【Bug已解决】AssertionError for FP8 Benchmarks 解决方案一、现象长什么样跑 FP8 推理 / 训练的 benchmark 脚本时中途直接抛AssertionErrorbenchmark 中断、无法出数。常见两种形态# 形态一断言输出 dtype 必须是 fp8但实际拿到 fp16 AssertionError out.dtype is not torch.float8_e4m3fn (got torch.float16) # 形态二断言 fp8 量化后的 scale 必须非零但出现了 0 AssertionError fp8 scale must be non-zero最小判据触发运行 FP8 benchmark量化 linear / matmul 现象中途 AssertionErrorbenchmark 不出结果 根因benchmark 的断言条件在真实 fp8 路径下不成立 影响benchmark 无法完成无法对比吞吐 / 精度最迷惑的是同样的断言在小规模玩具输入下能通过一旦 batch 拉大、或输入分布变化就炸。说明断言本身写得太脆没有覆盖真实 fp8 数值路径的全部情况。二、背景FP8 在 PyTorch 里有两种常见格式float8_e4m3fn动态范围小、精度高和float8_e5m2动态范围大、精度低。FP8 线性层通常走torch._scaled_mm或厂商内核流程是把输入xfp16/bf16按scale_x量化成 fp8权重w按scale_w量化成 fp8做 fp8 matmul得到 fp32 累加反量化回 fp16/bf16。benchmark 脚本为了验证 fp8 路径真的生效常写断言断言最终输出 dtype 是float8_*但真实 fp8 层反量化后输出是 fp16这个断言从根上就错断言scale ! 0但当某行输入全为 0 时该行的amax为 0scale amax / fp8_max也为 0这是合法情况断言误杀。于是 benchmark 的自检断言在合法 fp8 输入分布下失败benchmark 自己把自己打挂。这是断言条件与真实数值语义不符导致的 false positive。三、根因抽象成代码示意非照抄源码# benchmark 里的脆断言 def benchmark_fp8(x, w, sx, sw): x8 quantize_e4m3(x, sx) w8 quantize_e4m3(w, sw) out scaled_mm(x8, w8, sx, sw) # 内部反量化为 fp16 # BUG断言输出是 fp8但 scaled_mm 返回的是 fp16 assert out.dtype torch.float8_e4m3fn return out根因链条benchmark 想验证走了 fp8 路径选了输出 dtype 是 fp8作为判据但真实 fp8 线性层反量化后输出 fp16dtype 判据天然不成立另一个断言scale ! 0忽略了全零输入行 amax0 - scale0的合法情形断言在合法 fp8 分布下 false positivebenchmark 中断表现为AssertionError但根因是断言写错不是计算错。为什么玩具输入能过因为小输入里没有全零行、dtype 碰巧对掩盖了脆断言规模一大、分布一变就暴露。四、最小可运行复现用纯 Python 模拟全零行导致 scale0 触发误杀断言以及dtype 判据错误# repro_fp8_assert.py def quant_scale(amax, fp8_max448.0): # 真实 fp8amax0 时 scale 合法地为 0 return amax / fp8_max def benchmark_assert(amax_list): for amax in amax_list: scale quant_scale(amax) # BUG断言 scale 非零但全零行 amax0 - scale0 合法 if scale 0: raise AssertionError(fp8 scale must be non-zero (误杀合法全零行)) return True def main(): # 输入分布含一行全零amax0这是真实存在的 amax_list [1.2, 0.5, 0.0, 2.3] try: benchmark_assert(amax_list) print(benchmark 通过) except AssertionError as e: print(复现成功 -, e) if __name__ __main__: main()运行输出复现成功 - fp8 scale must be non-zero (误杀合法全零行)全零行的scale0是 fp8 量化的合法结果却被断言误杀正是真实 bug 的抽象。五、解决方案第一层最小直接修复最小且必须的一步把断言改成与真实 fp8 语义一致——不要断言输出 dtype 是 fp8而是断言中间确实走了 fp8 量化如断言x8.dtype是 fp8、或断言调用了scaled_mm不要把scale ! 0作为硬断言全零行的scale0合法应允许或在量化前对scale做max(scale, eps)以保证数值稳定但断言要放行 0。# fix_layer1.py def benchmark_fp8(x, w, sx, sw): x8 quantize_e4m3(x, sx) w8 quantize_e4m3(w, sw) # 正确断言量化后的张量确实是 fp8 格式验证走了 fp8 路径 assert x8.dtype in (torch.float8_e4m3fn, torch.float8_e5m2) assert w8.dtype in (torch.float8_e4m3fn, torch.float8_e5m2) out scaled_mm(x8, w8, sx, sw) # 反量化输出 fp16正常 # 不断言 out.dtype 是 fp8 return out这一层改动最小让断言描述真实发生的事。但它仍依赖手写的逐条断言容易再写错。六、解决方案第二层结构性改进把fp8 benchmark 的合法性校验收敛成一个校验器集中表达 fp8 数值规则避免散落的脆断言# fix_layer2.py from dataclasses import dataclass from typing import List FP8_DTYPES {float8_e4m3fn, float8_e5m2} dataclass(frozenTrue) class Fp8Check: quantized_dtypes: List[str] output_dtype: str scales: List[float] def validate_fp8_run(check: Fp8Check): # 规则1量化张量必须是 fp8 格式验证走了 fp8 路径 for dt in check.quantized_dtypes: assert dt in FP8_DTYPES, f量化张量 dtype{dt} 不是 fp8 # 规则2输出 dtype 允许 fp16/bf16反量化后不强制 fp8 assert check.output_dtype in FP8_DTYPES | {float16, bfloat16}, \ f输出 dtype{check.output_dtype} 非法 # 规则3scale 允许为 0全零行合法仅校验非负与有限 for s in check.scales: assert s 0 and s s, fscale{s} 非法需非负且有限 return True # 用法 validate_fp8_run(Fp8Check( quantized_dtypes[float8_e4m3fn, float8_e4m3fn], output_dtypefloat16, scales[1.2, 0.5, 0.0, 2.3], # 含合法的 0 ))要点校验集中表达fp8 路径 量化张量是 fp8、输出可反量化为 fp16、scale 非负有限全零行scale0被明确放行不再误杀用Fp8Check数据类描述一次运行新增 benchmark 只填数据不手写断言。七、解决方案第三层断言 / CI 守护写 pytest 验证合法 fp8 运行含全零行、含反量化输出不触发 AssertionError# test_fp8_benchmark.py import pytest FP8 {float8_e4m3fn, float8_e5m2} def validate(quant_dtypes, out_dtype, scales): for dt in quant_dtypes: assert dt in FP8 assert out_dtype in FP8 | {float16, bfloat16} for s in scales: assert s 0 and s s def test_zero_scale_allowed(): # 全零行 scale0 合法不应误杀 validate([float8_e4m3fn, float8_e4m3fn], float16, [1.2, 0.0]) def test_output_may_be_fp16(): # 反量化输出 fp16 是正常结果 validate([float8_e4m3fn], float16, [0.5]) def test_bad_quant_dtype_caught(): with pytest.raises(AssertionError): validate([float32], float16, [0.5]) # 没走 fp8 路径应被抓 def test_nan_scale_rejected(): with pytest.raises(AssertionError): validate([float8_e4m3fn], float16, [float(nan)])CI 一旦有人又把scale ! 0或out.dtype fp8写回断言相关测试立刻变红。八、排查清单FP8 benchmark 抛 AssertionError 时看清是哪条断言是 dtype 判据还是 scale 判据若是out.dtype fp8说明断言写错——fp8 层反量化输出本就是 fp16若是scale ! 0检查是否输入含全零行——那是合法 fp8 情况按第五 / 六节把断言对齐真实 fp8 语义用Fp8Check集中校验避免散落脆断言小输入能过、大输入才炸几乎可以断定是分布相关的误杀断言把第七节的 pytest 接进 CI守护合法 fp8 运行不出 AssertionError。九、小结FP8 benchmark 抛AssertionError根因是 benchmark 的自检断言与真实 fp8 数值语义不符要么错误断言输出 dtype 必须是 fp8实际反量化后是 fp16要么错误断言scale 必须非零全零行amax0时scale0是合法结果。断言在合法 fp8 分布下 false positivebenchmark 自我中断。三层层级第一层断言改为验证量化张量是 fp8、输出可反量化放行合法 scale0第二层用Fp8Check集中表达 fp8 数值规则消除散落脆断言第三层pytest 验证合法 fp8 运行含全零行、含 fp16 输出不触发断言锁进 CI。核心教训benchmark / 测试的断言必须描述真实成立的语义而非作者以为成立的前提。任何一个在小输入下侥幸通过、在大输入下炸掉的断言几乎都是断言写错了而非计算错了。

相关新闻