拓冰建站拓冰建站
首页 / 资讯中心 / 正文

torchtitan 手写 GPU Kernel 数值正确性验证:golden/reference/target 三方对拍方法论

torchtitan 手写 GPU Kernel 数值正确性验证golden/reference/target 三方对拍方法论【免费下载链接】torchtitanA PyTorch native platform for training generative AI models项目地址: https://gitcode.com/GitHub_Trending/to/torchtitan在 torchtitan 中替换默认算子如用 Triton kernel 加速 SwiGLU 激活、RMSNorm 或 RoPE时最大的风险不是性能而是数值正确性低精度 dtype、累加顺序、tile 边界都会悄悄引入偏差。本文基于仓库内置的 agent 技能 kernel-numerics-verifier完整讲解其“三实现三方对拍”验证框架如何构造 golden/reference/target 三套实现、用什么样的判据correctness gate判定 kernel 通过、在 reference 本身不够精确时如何做 conditioning 分析以及一套 9 步工作流与结果报告规范。文中还会结合 torchtitan 仓库中真实落地的 Triton kernel 及其数值测试如 Qwen3.5 OffsetRMSNorm 的 Triton override 与对应测试说明该技能的各项参数在工程代码里如何取值、形状与对抗样本怎么选使读者既能掌握方法论本身也能在仓库中找到可复现的实现证据。为什么需要三方对拍torchtitan 的 kernel 扩展场景torchtitan 通过 override 机制允许把模型中的任意Configurable组件替换为自定义实现其中典型用途就是硬件相关的 Triton/CUDA kernel——override 机制文档将hardware-specific kernels列为第一类动机并明确仓库内提供了三个真实示例fused_swiglu.py用融合的 Triton kernel 替换 SwiGLU 的 SiLU 与逐元素乘法前向 反向两个 kernel注册为torch.library.custom_opoffset_rmsnorm.py为 Qwen3.5 的OffsetRMSNorm提供融合 Triton 前向/反向 kernel注册为torch.library.triton_ophelion_rope.py用 Helion 融合 kernel 替换CosSinRoPE。这些 kernel 都是低精度输入 FP32 累加的手写实现与 eager 参考实现之间天然存在数值差。skill 文档SKILL.md的 description 也划清了适用边界本技能用于验证单个手写 kernel 的前向/反向数值正确性如果要定位的是整个模型的未知数值漂移应改用numerics_debugging技能基于DebugMode的逐算子激活捕获与对比。两者是互补关系本文聚焦前者。核心方法golden、reference、target 三个实现skill 要求对同一算子准备三个实现每个承担不同的尺子角色实现定义作用goldeneager 的 FP64 实现FP64 不受支持时才退化为 FP32近似真值的高精度 oraclereferenceeager 实现且 dtype 与 kernel 的输入/输出 dtype 完全相同与 target 处于同一精度起跑线的基线target手写的 Triton 或 CUDA kerneldtype 与reference相同被验证对象其中最关键的一个细节是输入生成方式只生成一份低精度输入然后把这份完全相同的值无损失地提升到 FP64 给golden使用。这样测量到的是纯 kernel 算术误差不会混入输入量化误差。这一点在 torchtitan 的真实测试中有严格体现。test_triton_offset_rmsnorm_override.py 中先按目标 dtypebf16/fp16/fp32生成一份input_data/weight_data/grad_output_data然后三路各自clone()golden一路.double()后在 FP64 下做 eager RMSNorm 并 autograd 求梯度reference一路保持原 dtype 走 FP32 中间计算的 eager 公式target一路调用triton_offset_rms_norm。对应的两个 oracle 实现如下def _offset_rms_norm_reference( input: torch.Tensor, weight: torch.Tensor, ) - torch.Tensor: input_dtype input.dtype input_fp32 input.float() inverse_rms torch.rsqrt(input_fp32.square().mean(-1, keepdimTrue) _EPS) return ((1.0 weight.float()) * input_fp32 * inverse_rms).to(input_dtype) def _offset_rms_norm_golden( input: torch.Tensor, weight: torch.Tensor, ) - torch.Tensor: input_fp64 input.double() inverse_rms torch.rsqrt(input_fp64.square().mean(-1, keepdimTrue) _EPS) return (1.0 weight.double()) * input_fp64 * inverse_rms注意reference与targetTriton kernel 内部同样是 load 后.to(tl.float32)累加走同一套 dtype 策略因此二者之差才能干净地归因于 kernel 本身的算术实现而不是参考实现天生更差。正确性判据Correctness Gate对每一个输出张量和梯度张量都在 golden 的 dtype即 FP64下计算以下五个量reference_error max(abs(reference - golden)) target_error max(abs(target - golden)) rounding_floor max(abs(golden.to(target_dtype).to(golden_dtype) - golden)) absolute_floor max(project_atol, rounding_floor) threshold fudge_factor * reference_error absolute_floor各项含义reference_erroreager 参考实现相对 oracle 的偏差代表同精度下能达到的最好水平target_error被测 kernel 相对 oracle 的偏差rounding_floor把 golden 值先舍入到 target dtype 再转回 FP64 的误差即输出 dtype 的舍入下限——kernel 结果不可能比这个更接近 golden它构成了绝对底线threshold判据阈值。核心思想是kernel 误差至多是参考实现的fudge_factor倍再叠加舍入底线。当 reference 足够精确时判据为target_error threshold关于两个自由参数skill 给出了明确的纪律性要求若算子已有官方容差如 operator tolerance优先沿用否则从fudge_factor 2.0和project_atol 0起步必须在运行 target 之前预先声明这两个值禁止为了让失败的 kernel 通过而调大它们。仓库测试与该规范逐字吻合test_triton_offset_rmsnorm_override.py 顶部就固定了_EPS 1e-6 _FUDGE_FACTOR 2.0 _PROJECT_ATOL 0.0 _MAX_REFERENCE_RELATIVE_ERROR 0.1而实际的门控逻辑_assert_matches_golden同文件 L47-L92完整实现了上述公式先断言 shape、dtype、NaN 位置、正/负 Inf 位置一致再在 FP64 下算reference_error、target_error、rounding_floor组成threshold _FUDGE_FACTOR * reference_error absolute_floor并断言target_error threshold。这与 skill 工作流第 3 步数值容差之前先比结构的要求一一对应。先检查条件数再谈通过/失败Conditioning Checkgate 有一个隐含前提reference 本身相对 golden 的误差足够小。skill 要求在套用 gate 之前先度量这一点golden_scale max(abs(golden)) reference_relative_error reference_error / max(golden_scale, absolute_floor)预先声明一个可接受的 reference 相对误差若没有现成标准用 10% 作为诊断触发阈值测试代码中的_MAX_REFERENCE_RELATIVE_ERROR 0.1正是这个默认值并在超限时给出明确报错reference is too inaccurate to gate the target一旦 reference 相对误差超过该阈值就不得再给出 PASS/FAIL 结论——因为此时允许的误差已经大到无法区分kernel 有 bug和问题本身数值敏感。归约类算子逐元素条件数诊断对于归约reduction类输出skill 要求定位到贡献最大误差的那个元素并把它的 FP64 逐项贡献抓出来condition_number sum(abs(per_term)) / abs(sum(per_term))条件数大说明相消cancellation会放大项级小误差。但要注意条件数大本身不能证明 target 是对的它只解释了误差为什么这么大。gate 失效时的三步替代分析当 gate 无效reference 相对误差超阈值时skill 规定执行以下证据收集流程而不是硬套 2 倍阈值对比两种合法的 FP32 归约顺序——直接测量对累加顺序的敏感度跨多个归约长度 T 和多个随机种子重复误差随sqrt(T)大致增长说明是零均值数值噪声随T线性增长说明存在系统性偏差。同时记录每个T下的条件数对比 reference 与 target 的误差分布和 bias。target 特有的偏置、异常缩放或语义不匹配直接判 FAIL。结论判定规则若两个实现都遵循同一套条件数驱动的噪声包络报告PASS WITH LIMITATIONS若证据无法把数值噪声与 kernel bug 区分开报告INCONCLUSIVE——宁可不下结论也不给出虚假的通过。逐元素 kernel 的 ULP 要求对简单的逐元素elementwisekernelskill 额外要求结果必须是正确舍入correctly rounded或误差在 1~2 ULP 以内。超过这个倍数就必须拿出证据——例如已有被接受的实现在代表性输入上验证过该放大因子是合理的——而不是随口调宽阈值。九步工作流skill 给出的完整验证流程如下与仓库测试的组织方式可直接对照读代码。读 kernel 源码、eager 公式、调用点和现有测试。记录公式、中间 dtype、累加 dtype、输出 dtype、支持的形状、masking 行为、NaN/Inf 行为。构造输入。为reference和target生成完全相同的低精度输入golden输入由这些值无损失升格而来。先比结构再比数值。比较 shape、dtype、NaN 位置、正/负 Inf 位置任何一处不一致都直接判失败。检查 reference 相对误差。只有 reference 足够精确才套用 gate否则走上面的 conditioning 分析。对每一个前向输出应用上述决策流程不能只抽查一个。反向验证。三个实现使用同一个随机上游梯度然后逐个检查每个输入/参数梯度。形状覆盖。测生产形状、最小合法形状以及紧贴 kernel tile 大小上下各一维的尺寸并加入与该算子相关的对抗值如相消、极端 logits、近零方差。非确定性 kernel重复运行相同输入要求每一次结果都通过同一 gate。先诊断后放宽。失败时按顺序排查indexing 与 mask → 累加 dtype → 过早 cast → 归约顺序 → 数值稳定性处理 → 近似函数 → 反向公式确认根因之前不得改容差。结果报告规范skill 要求为前向和每个梯度输出一个紧凑表格TensorRef/goldenConditionTarget errorThresholdGate validVerdict最终只给一个总判定PASS、PASS WITH LIMITATIONS、FAIL或INCONCLUSIVE。同时必须说明测过的 shapes、dtypes 和未覆盖的情况不得从单一形状或只测前向就宣称通用正确数值验证通过之后才做性能 benchmark——性能测试永远排在正确性之后。仓库实例OffsetRMSNorm Triton kernel 的完整落地tests/unit_tests/gpu/test_triton_offset_rmsnorm_override.py 是本 skill 在 torchtitan 中的一次近乎教科书式的落地可以逐项对照上文被测 kerneltorchtitan/overrides/offset_rmsnorm.py 为 Qwen3.5 的OffsetRMSNorm计算(1 weight) * rmsnorm(input)。实现要点对应工作流第 1 步记录公式与 dtype前向 kernel 每行一个 programload 后.to(tl.float32)累加方差用tl.sum(x * x) / num_cols再tl.rsqrt(variance eps)L48-L71block_size triton.next_power_of_2(num_cols)上限_MAX_BLOCK_SIZE 65536num_warps按 block 大小分档 4/8/16L35-L45——tile 边界就是第 7 步要重点覆盖的形状边界前向额外输出inverse_rms供反向复用反向 kernel 计算输入梯度和分块归约的权重梯度通过torch.library.triton_op注册L171并register_autograd挂接反向保证torch.compile可追踪——这正是 override 机制文档Custom kernels and torch.compile一节要求的注册配方仅在input.is_cuda且 dtype 属于(float16, bfloat16, float32)时走 Triton 路径否则回退 eager 实现L310-L316。形状选择紧贴 tile 边界的上下各一维test_forward_and_backward_against_golden的用例表L177-L191精确演示了工作流第 7 步dimensions immediately below and above relevant kernel tile sizescases ( ((3, 255), torch.float32, 1), ((3, 256), torch.bfloat16, 2), ((3, 257), torch.bfloat16, 3), ((3, 257), torch.float16, 10), ((8, 6, 256), torch.bfloat16, 4), ((8, 4095), torch.bfloat16, 5), ((8, 4096), torch.bfloat16, 6), ((8, 4097), torch.bfloat16, 7), ((8, 5120), torch.bfloat16, 8), )255/256/257 与 4095/4096/4097 三组正好卡在next_power_of_2的跳变点上257 → 512 块、4097 → 8192 块覆盖 mask 分支mask col_offsets num_cols在整块对齐与非对齐两种情况下的行为。5120 则是一个生产级维度。对抗样本近零方差与零方差RMSNorm 在方差趋近 0 时对rsqrt与除法最敏感属于 skill 所说的near-zero variance对抗值def test_near_zero_variance(self): self._run_case((8, 5120), torch.bfloat16, 9, scale1e-5) def test_zero_variance(self): self._run_case((8, 5120), torch.bfloat16, 11, scale0.0)反向验证单一上游梯度、逐梯度检查_run_case为三个实现生成同一份grad_output_data分别调用torch.autograd.grad得到grad_input/grad_weight然后对输出和每个梯度张量逐一套用_assert_matches_golden——与第 5、6 步完全一致。附加契约测试数值之外的两项工程验证也值得注意test_custom_op_contractL199-L227用torch.library.opcheck检查前向/反向 op 的 schema、faketensormeta kernel、autograd 注册一致性——对应 override 文档建议把 opcheck 作为 override 包自带的单元测试test_torch_compileL229-L264以fullgraphTrue编译后对比 eager 的输出与梯度确认 kernel 在torch.compile下不 graph-break 且数值不变。同目录下的 test_fused_swiglu.py 则展示了验证的另一半——checkpoint 布局兼容性fused 实现保存逻辑布局w1.weight/w3.weight并与 native 实现、HF adapter 互相load_state_dict后逐位一致且strict加载仍能报告真正缺失的键。对 SwiGLU 这类非归约的逐元素 kernel其前向/反向 kernel 同样遵循FP32 累加、mask 对齐 tile的写法fused_swiglu.py L44-L165。与模型级数值调试技能的分工skill 的 frontmatter 明确了两者的分工原文Use for Triton or CUDA kernel correctness; use numerics_debugging to locate an unknown whole-model divergence.kernel-numerics-verifier本文主题单算子粒度三方对拍 gate 判据回答这个 kernel 写得对不对numerics_debugging模型级粒度用torch.utils._debug_mode.DebugMode在指定 step 捕获逐算子激活activation_tracer.py再用compare_numerics.py对两次运行做 diff 生成 HTML 报告回答两次本应一致的 run 在哪里开始漂移。实践中合理的路径是训练层面发现漂移 → 用 numerics_debugging 定位到可疑算子 → 对该算子切换为 kernel-numerics-verifier 的三方对拍做定论。小结torchtitan 的 kernel 数值验证技能把手写低精度 kernel 对不对这一模糊问题收敛为一套可执行、可复判的程序三实现分工明确——FP64 oracle 定真值、同 dtype eager 定基线、kernel 定对象且 golden 输入必须由低精度输入无损失升格得到gate 公式threshold fudge_factor * reference_error absolute_floor把相对参考实现的误差放大倍数和输出 dtype 舍入底线分开计量两个系数必须预先声明默认 2.0 / 0不许事后放水在 gate 前提失效时用 reference 相对误差10% 诊断阈值 逐元素条件数 sqrt(T)/T增长判别把噪声和bug区分开允许以PASS WITH LIMITATIONS或INCONCLUSIVE收场9 步工作流强制覆盖 tile 边界形状、对抗值、全部前向输出与全部梯度、非确定性重放结果以单表 单一 verdict 交付性能测量严格排在数值通过之后。OffsetRMSNorm 的 Triton 实现与测试证明了这套规范在 torchtitan 中的可执行性常量、gate、tile 边界形状、对抗方差、opcheck 与torch.compile契约测试都能在源码中逐条找到对应读者可直接以它为模板为自家 Triton/CUDA kernel 编写同级别的数值验证。【免费下载链接】torchtitanA PyTorch native platform for training generative AI models项目地址: https://gitcode.com/GitHub_Trending/to/torchtitan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门