# 多卡通算融合算子精度校验 > 面向多卡(NPU)环境下"通信 + 计算"融合算子的精度可信度验证。 > 核心思路:**把融合算子拆成通信与计算两部分分别独立校验**,再**与 PyTorch 分布式官方实现做整体基准对比**,从而把"整体跑通"细化成"每一环都对"。 ## 1 背景与问题 融合算子在 `forward()` 里把跨卡数据搬运(collective)与本地计算(gemm / 逐元素 / 重排)压进一个内核。这种写法性能好,但给精度判定带来困难: 1. **错误环节判断困难**:整体输出与 golden 不一致时,无法直接判断是通信搬错了数据、还是计算得了错误结果。 2. **single-pass 无法复现**:多卡可能出现进程间先后的错误,单卡测试掩盖了通信路径问题。 3. **基准本身要可信**:PyTorch 分布式实现是我们的对照基准,基准本身若配置错误(world_size、分组、后端),则所有对比失真。 因此,精度验证不应只做"最终结果比对"这一步,而应分层: ``` 第 1 层:通信路径校验 —— 只关心数据是否按预期到达、拼装、归约 第 2 层:单卡计算语义校验 —— 只关心 NPU 内核的数值结果是否正确 第 3 层:融合整体基准对比 —— 在真实多卡上,与 PyTorch 分布式实现全量比对 ``` 三层都通过,才能给融合算子的精度结论。 ## 2 总体验证框架 ### 2.1 三个执行入口 | 层级 | 校验对象 | 运行方式 | 判定目标 | |------|----------|----------|----------| | 通信校验 | 通信子模块(单独隔离) | 单核或多核,只发不算是核心 | 搬移/拼装/归约结果正确 | | 计算校验 | 本地计算内核 | 单卡,喂固定输入 | 数值语义正确(含 NaN/边界) | | 整体对比 | 融合算子 | 真实多卡分布式 | 与 PyTorch 分布式基准一致 | ### 2.2 通用校验工具 建议在工程里固化下列工具,减少手工比对: - **Golden 生成**:用 PyTorch 分布式官方 API 作为真值来源(`dist.all_gather` / `all_reduce` / `reduce_scatter` / `all_to_all_single` 等)。 - **固定 seed**:所有随机输入用固定 seed,保证可复现。 - **宽松 + 严格双阈值**:数值对比用 `allclose`(默认相对/绝对容差),另外单独做**逐位或极严格**检查的前提是输入恰好可用二进制精确表示的情形。 ## 3 精度对比与误差分析 三层校验最终都落到同一个动作:把"计算结果"与"参考结果(Golden)"做数值对比。本节定义这一动作的统一实现、评判标准与注意事项。 ### 3.1 获取 Golden Golden 是判定真值的来源,优先级由高到低: 1. **等价 Torch 算子**在 CPU / GPU 上以高精度(`fp32` / `fp64`)计算得到的结果; 2. **相同 Triton 算子**在 CPU / GPU 上的计算结果(用于核对实现语义一致性)。 > 说明:Golden 应尽量用更高精度计算,避免把参考实现自身舍入误差算进容差。若目标 dtype 为 `fp16`/`bf16`,建议用 `fp32` 计算 golden 再转回目标 dtype。 ### 3.2 对比判断入口 > 分布式环境下,融合算子不再只是单卡上的"本地计算",而是**跨卡先搬运数据、再在本地完成重排/计算**(通信 + 计算融合)。以下以 **HCCL Reverse All-to-All** 融合算子为例:每个 rank 持有本地分片,跨卡通过对称内存把所有 rank 的分片搬运、重排成完整结果。示例来自 `04-ascend-reverse-all2all.py`。 分布式 Reverse All-to-All 示例(融合通信 + 计算): https://gitcode.com/Ascend/Triton-distributed-ascend/blob/master/tutorials/ascend/04-ascend-reverse-all2all/04-ascend-reverse-all2all.py > 说明(示例要点): > - 用 `dist.init_process_group(backend="hccl")` 初始化进程组,并通过 ACL 对称内存(`aclshmem_*`)完成跨卡数据搬运; > - 融合内核 `kernel_hccl_reverse_a2a_pipelined` 在单个 `@triton.jit` 内核内同时完成跨卡搬运与本地重排(Producer/Consumer 阶段); > - Golden 用 PyTorch 分布式官方实现 `torch_reverse_a2a`(`dist.all_to_all_single`)构造,再与融合算子输出做 `assert_close` 比对; > - 精度对比通过 `dist.all_gather` 收集各 rank 的通过标志,全部通过才判通过。 对应运行命令(多卡启动,`LOCAL_RANK` 由 torchrun 自动注入): ```bash torchrun --nproc_per_node=2 --master_port=29500 04-ascend-reverse-all2all.py ``` ### 3.3 精度对比函数(统一入口) ```python def compare_precision(cal, ref, rtol=1e-3, atol=1e-3): """ 精度对比函数:根据数据类型选择合适的比对策略。 参数: cal: 计算结果 ref: 参考结果 rtol: 相对误差容限 atol: 绝对误差容限 异常: AssertionError: 精度不达标时抛出 """ assert cal.dtype == ref.dtype, f"dtype mismatch: {cal.dtype} vs {ref.dtype}" tensor_dtype = cal.dtype if tensor_dtype == torch.float16: torch.testing.assert_close(ref, cal, rtol=rtol, atol=atol, equal_nan=True) elif tensor_dtype == torch.bfloat16: torch.testing.assert_close(ref, cal, rtol=5e-3, atol=5e-3, equal_nan=True) elif tensor_dtype == torch.float32: torch.testing.assert_close(ref, cal, rtol=1e-5, atol=1e-5, equal_nan=True) elif tensor_dtype in [torch.int64, torch.int32, torch.int16, torch.int8]: assert torch.equal(cal, ref), f"Integer tensors are not equal for dtype {tensor_dtype}" elif tensor_dtype == torch.bool: assert torch.equal(cal, ref), "Boolean tensors are not equal" else: raise ValueError(f"Unsupported tensor dtype: {tensor_dtype}") print(f"dtype: {tensor_dtype} — Precision check passed.") ``` ### 3.4 判定标准 - `torch.testing.assert_close` / `torch.equal` **不抛出异常 → 通过**,否则不通过。 - `torch.testing.assert_close`:张量在指定容差内近似相等即通过,否则抛出 `AssertionError`。 - `torch.equal`:形状相同且所有元素二进制级绝对相等才返回 `True`。 - `torch.testing.assert_close` 的内部判据: ``` |cal - ref| <= atol + rtol * |ref| ``` 即绝对误差必须落在"绝对容限 + 相对容限 × 参考值幅度"构成的动态边界内。当 `ref` 接近 0 时,`atol` 起主导作用;当 `ref` 较大时,`rtol` 起主导作用。 按数据类型推荐容限: | 数据类型 | rtol | atol | 说明 | |----------|------|------|------| | float32 | 1e-5 | 1e-5 | 严格 | | float16 | 1e-3 | 1e-3 | 精度较低,适当放宽 | | bfloat16 | 5e-3 | 5e-3 | 精度较低,适当放宽 | | int8/16/32/64 | — | — | 必须完全一致(`torch.equal`) | | bool | — | — | 必须完全一致(`torch.equal`) | > 说明:上述容限为工程默认值。熔融合算子因叠加通信与计算误差,可比单算子略微放宽 `fp16`/`bf16`,但**必须明确写入测试文档**,不能静默放宽。 ### 3.5 注意事项 - **NaN / Inf 处理**:`equal_nan=True` 会把 NaN 视为相等。若需要严格检测 NaN 差异,应设为 `False`。是否传播 NaN 属于算子语义,测试要对应语义选择该开关。 - **整数类型**:`int` / `bool` 不允许任何误差,必须严格一致。 - **设备对齐**:在跨设备对比时,务必把 `cal` 与 `ref` 搬到**同一设备(如 CPU)**再比较,避免底层表示差异造成误判。注意 `triton_cal.cpu()` 应放在比较之前。 - **dtype 对齐**:比较前先断言 `cal.dtype == ref.dtype`,dtype 不一致直接失败,报错信息要明确。 ## 4 通信路径校验 通信校验的目标是**证明跨卡的数据移动本身正确**,与计算无关。因此应把通信行为从融合内核中隔离出来,单独构造、单独比对。 ### 4.1 隔离方式 - 构造一个"只通信、不计算"的测试内核:它做的事情就是从对端 rank 的对称内存 `load`/`store`、拼装 shard、或做纯 `identity` 的归约(不做 gemm / 逐元素变换)。 - 若算子结构上难以单独导出通信内核,可用"单位算"替身:把计算部分替换成 `x = x`,保留通信与重排结构,从而单独暴露通信行为。 ### 4.2 校验点 对每类 collective,锁定其关键不变量: | 通信原语 | 必查不变量 | |----------|-----------| | `all_gather` | 每个 rank 收到的拼接结果是各 rank 本地分片按 rank 序拼接;长度 = world_size × 本地长度 | | `all_reduce` | 各 rank 结果一致;数值 = 各 rank 对应元素之和(对 SUM) | | `reduce_scatter` | rank i 持有的分片 = 各 rank 对应分片归约结果 | | `all_to_all_single` | rank i 从 rank j 收到的块 = rank j 想要发送给 i 的块;块内容、块序、块大小一一对应 | ### 4.3 判定 - 对每个 rank 独立比较通信输出与 PyTorch 对应 collective 的 golden,**每个 rank 都通过**才算通过。 - 通信层不要用宽松的逐元素替换掩盖错误路径;优先用**足够有区分度的输入**(例如每个 rank 用不同数值的块),以便一旦错位立刻暴露。 ### 4.4 关注点 - **rank 序 / 分块粒度**:最容易出错的是"谁发给了谁、拼到哪个位置"。 - **对称内存布局**:offset / stride / shard 形状一旦与 launcher 参数不匹配,会出现跨 rank 错位但进程不报错的情况。 - **同步**:通信正确还依赖 barrier / wait / notify 顺序;通信校验同时观察是否出现悬挂或数据未就绪。 ## 5 计算语义校验 计算校验的目标是**证明本地数值计算正确**,与通信无关。在单卡上直接对内核给出确定输入做比对,可快速定位数值错误。 ### 5.1 隔离方式 - 在单卡环境下,用"身份通信"替代真实通信:数据不跨卡,直接用本地内存,使内核的计算路径可被单独测试。 - 保持与融合内核**相同的 tile / block / num_stages / 布局参数**,避免"单测通过、融合失败"的参数漂移。 ### 5.2 校验点 - **数值语义**:与 `torch` 对应操作(`matmul` / `elementwise`)在相同输入下 allclose。 - **NaN / Inf 传播**:若算子语义要求 NaN 传播,用含 NaN 输入确认传播行为符合预期。 - **边界与尾部**:关注非整除的分块、mask / 边界填充是否正确处理最后一块。 - **精度等级**:`fp16`/`bf16` 需要说明容差选择,`fp32` 可收紧。 ### 5.3 判定 - 计算层单卡全部通过后,才进入整体对比,避免把计算 bug 带到多卡排查。 ## 6 融合整体基准对比 在真卡多卡环境下执行融合算子,与 PyTorch 分布式基准做全量比对。这是最终精度结论的唯一裁决入口。 ### 6.1 基准配置要点 - **基准实现**:用 PyTorch 分布式官方 API(`dist.init_process_group(backend="hccl")` + 对应 collective)搭出与算子输入输出完全一致的参照实现。 - **world_size / rank 对齐**:基准与融合算子必须使用相同的进程组规模与分组;任何不一致都会让对比失真。 - **输入完全一致**:对比时给基准与融合算子**同一份输入张量**(或按 rank 复制到独立内存但内容逐位相同),消除输入差异。 - **布局与 dtype 一致**:输入/输出内存布局(连续/非连续、NCHW/NHWC)、dtype 都必须一致。 ### 6.2 对比流程 ``` 1. 初始化多卡进程组,barrier 确保各 rank 就绪 2. 各 rank 用固定 seed 生成输入 3. 分别运行 PyTorch 基准 与 融合算子 4. 各 rank 独立比较输出 allclose 5. 汇总:收集各 rank 的通过标志,全部通过才判通过 ``` ### 6.3 判定 - **逐 rank 独立 allclose**:每个 rank 的输出都与该 rank 的 golden 比对。 - **跨 rank 一致性检查(可选)**:对要求各 rank 输出一致的算子(如 all_reduce),额外确认各 rank 结果一致。 - **判定标准**:只有所有 rank 全部 allclose,且无悬挂、无 crash,才算整体精度通过。 ## 7 常见问题排查 | 现象 | 可能原因 | 建议动作 | |------|----------|----------| | 整体不一致,但通信层通过 | 计算语义 bug | 回到第 2 层排查计算 | | 整体不一致,但计算层通过 | 通信搬移/拼装错误 | 回到第 1 层排查通信 | | 单 rank 通过、多 rank 失败 | rank 序/分块/对称内存布局问题 | 重点查第 1 层 4.4 | | 结果一致但极值处差 | 精度/容差选择不当 | 说明数值精度等级,调整容差 | | 进程悬挂 | 同步(barrier/wait/notify)错序 | 检查同步顺序,必要时加 barrier | ## 8 落地清单 - [ ] 固定 seed、可复现的输入生成器 - [ ] 通信隔离测试(identity 替身,逐 collective 校验) - [ ] 单卡计算语义测试(含 NaN / 边界) - [ ] 多卡整体 vs PyTorch 分布式基准对比 - [ ] 每个 rank 的 allclose 通过标志收集与汇总 - [ ] 精度等级与容差说明写入文档 ## 9 总结 精度验证的正确打开方式是**分层**: 1. **通信与计算分开**,各用对应的最小单元证伪/证实正确性; 2. **最后在真实多卡上**,用 PyTorch 分布式实现做权威基准对比。 分层可以把"整体对不上"快速归因到通信或计算,避免在多卡、难复现的环境里大海捞针。