多卡通算融合算子精度校验

面向多卡(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_a2adist.all_to_all_single)构造,再与融合算子输出做 assert_close 比对;

  • 精度对比通过 dist.all_gather 收集各 rank 的通过标志,全部通过才判通过。

对应运行命令(多卡启动,LOCAL_RANK 由 torchrun 自动注入):

torchrun --nproc_per_node=2 --master_port=29500 04-ascend-reverse-all2all.py

3.3 精度对比函数(统一入口)

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 不允许任何误差,必须严格一致。

  • 设备对齐:在跨设备对比时,务必把 calref 搬到**同一设备(如 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 分布式实现做权威基准对比。

分层可以把"整体对不上"快速归因到通信或计算,避免在多卡、难复现的环境里大海捞针。