1. 从一次训练崩溃说起:为什么 Bitwise 一致这么难
做过强化学习训练的人大概都经历过这种场景:同一个模型、同一份数据、同一组超参,跑两次 loss 曲线就是对不上,rollout 阶段采样出来的轨迹和训练阶段重算的 logprob 差了一点点,单步看不出来,累积几百步之后策略直接跑偏。你盯着屏幕怀疑人生,怀疑随机种子,怀疑数据加载顺序,最后发现是 kernel 层面的浮点累加顺序不一样。
这就是Bitwise 一致(逐位一致)要解决的问题。它不是"误差很小",而是"二进制层面完全相同"。在 RL 训练里,rollout 和 training 是两个阶段,前者用推理引擎跑采样,后者用训练框架算梯度。如果这两个阶段对同一个 token 序列算出来的 logprob 在 bit 级别有差异,那么重要性采样比(importance ratio)就会偏离 1,PPO 这类算法的 clip 机制会被这些虚假差异触发,训练稳定性直接崩掉。
我这次折腾的项目,核心目标就一句话:在 AMD 平台上,用 RL-Kernel 配合 vime 框架,让训练阶段和 rollout 阶段的前向计算做到 bitwise 一致。关键词拆开看:RL-Kernel 是算子层的东西,vime 是 RL 训练框架,AMD 是硬件平台,ROCm 是软件栈。这四个东西凑一起,坑比想象中多得多。
先说清楚适合谁看。如果你正在做 RLHF、GRPO、PPO 这类需要 rollout 和 training 分离的训练任务,并且用的是 AMD 的卡(比如 MI 系列或者消费级的 Radeon),那这篇内容基本就是给你写的。如果你只是跑跑 SFT,或者用的是单阶段训练,那 bitwise 一致这个问题对你影响没那么大,但了解一下 kernel 层面的数值行为也没坏处。小白读者也不用慌,我会把浮点、kernel、ROCm 这些概念用生活化的方式讲清楚,保证你能跟上。
2. 整体设计思路:为什么要在 AMD 上死磕这件事
2.1 问题的本质:rollout 和 training 是两套计算路径
先把这个问题的根源讲透。RL 训练里,rollout 阶段通常用推理引擎(比如 vLLM 那套思路的引擎)来生成轨迹,追求的是吞吐;training 阶段用训练框架(PyTorch + 各种并行策略)来算梯度,追求的是精度和显存效率。这两个阶段的算子实现往往是两套完全不同的代码。
推理引擎为了快,会把 attention 做成 paged attention,把 KV cache 管理得花里胡哨,矩阵乘可能用特定的 fused kernel;训练框架为了稳,用的是另一套 attention 实现,可能还带 flash attention 的变体。这两套实现在数学上等价,但在浮点运算上不等价——因为浮点加法不满足结合律。
举个生活化的例子:你要把 100 个小数加起来。从左往右加,和从右往左加,或者先分组加再合并,结果在实数域上一样,但在浮点数域上可能差最后几位。GPU kernel 的并行归约(reduction)就是典型的分组加法,不同的线程块划分、不同的 warp 调度,归约顺序就不同,最后那几位就不一样。
注意:bitwise 一致不是要求所有算子都一模一样,而是要求影响 logprob 计算的那条关键路径上,rollout 和 training 的浮点运算顺序完全一致。其他不影响 logprob 的部分(比如 KV cache 的存储布局)可以不同。
2.2 为什么选 RL-Kernel 作为切入点
RL-Kernel 这个思路的核心是:把 rollout 和 training 共用的那部分计算,抽成同一套 kernel。与其让两个框架各自实现一遍然后祈祷它们数值一致,不如让它们调用同一个底层算子。
这个设计的好处很直接:
- 消除实现差异:同一个 kernel,同样的线程划分,同样的归约顺序,输入相同则输出 bitwise 相同。
- 便于验证:只需要验证一套 kernel 的正确性,不用交叉验证两套实现。
- 性能可控:kernel 层面可以做针对性的优化,比如针对 AMD 的 wavefront 大小(64)来调整线程块。
代价也很明显:这套 kernel 必须同时满足推理的高吞吐和训练的高精度需求,设计上要做取舍。比如推理可以用 fp16 甚至 fp8 来加速,但训练如果也用低精度,梯度就会炸。所以 RL-Kernel 通常会在关键路径上用 fp32 累加,即使输入是低精度。
2.3 AMD 平台带来的额外变量
为什么在 AMD 上这件事更难?因为 ROCm 生态和 CUDA 生态有差异,很多 kernel 的默认行为不一样。
第一,wavefront 大小不同。NVIDIA 的 warp 是 32 线程,AMD 的 wavefront 是 64 线程(RDNA 架构有 32 和 64 两种模式)。这意味着同样的归约算法,在 AMD 上的分组方式和 NVIDIA 不同,浮点累加顺序自然不同。如果你直接拿 CUDA 的 kernel 翻译成 HIP,数值行为可能就变了。
第二,数学库实现不同。rocBLAS、rocFFT、MIOpen 这些库和 cuBLAS、cuDNN 的实现细节不一样。比如 exp、log、tanh 这些超越函数的近似算法可能不同,最后几位就有差异。
第三,编译器优化行为不同。AMD 的编译器(比如 hipcc 背后的 LLVM)在做 fast-math 优化时,对浮点重排的策略和 NVCC 不一样。-ffast-math这种选项在两边打开,结果可能差很多。
所以整个项目的设计思路是:先用 RL-Kernel 统一关键路径的算子实现,再针对 AMD 的硬件特性做适配,最后用 bitwise 对比工具验证 rollout 和 training 的输出。这三步缺一不可。
3. 核心细节解析:Bitwise 一致到底卡在哪些环节
3.1 浮点累加顺序:最隐蔽的杀手
前面提到浮点加法不满足结合律,这里展开讲。假设你要算a + b + c + d,四种可能的顺序:
((a+b)+c)+d (a+(b+c))+d (a+b)+(c+d) a+((b+c)+d)在实数域上结果一样,在 IEEE 754 浮点下可能四个结果都不同。GPU 上的归约操作,比如 softmax 里的求和、attention 里的加权平均、layer norm 里的均值方差,全都是归约。归约的并行实现通常是树形的:先把数据分到多个线程,每个线程局部求和,然后再跨线程合并。线程数、每个线程处理多少元素、合并的顺序,任何一个变了,结果就变。
在 RL 训练里,最敏感的归约是logprob 的计算。logprob 通常是对 logits 做 log_softmax 然后取对应 token 的值。log_softmax 里有 max 归约和 sum 归约,这两个归约的顺序直接决定 logprob 的最后几位。
我实测过一个案例:同一个模型,rollout 用推理引擎算出的 logprob 和 training 用训练框架算出的 logprob,在 fp16 下差异能到 1e-3 量级,在 fp32 下能到 1e-6。听起来很小,但 PPO 的 ratio 是exp(logprob_new - logprob_old),1e-6 的差异经过 exp 放大,再累积几千个 token,ratio 就能偏离 1 好几个百分点。
提示:判断你的训练是否受这个问题影响,最简单的办法是打印 rollout 和 training 对同一批数据的 logprob,算一下 max abs diff。如果大于 1e-5,基本可以确定是 kernel 层面的数值不一致。
3.2 RL-Kernel 的关键设计:统一归约树
RL-Kernel 解决这个问题的核心手段是固定归约树。具体做法:
- 对于所有影响 logprob 的归约操作,强制使用同一套线程划分和归约顺序。
- 归约的中间结果用 fp32 存储,即使输入是 fp16/bf16。
- 禁用编译器的 fast-math 重排,或者用
#pragma显式指定归约顺序。
这里有个细节:AMD 的 wavefront 是 64 线程,做 warp-level reduction 的时候,shuffle 操作的语义和 NVIDIA 不同。RL-Kernel 在 AMD 上的实现需要针对 wavefront 大小做特化。比如一个 64 线程的 wavefront 做归约,通常是先做 32 对 32 的 shuffle,再做 16 对 16,以此类推。这个顺序必须固定下来,不能依赖编译器的自动优化。
我踩过的一个坑:一开始用__shfl_down做归约,在 NVIDIA 上没问题,移植到 AMD 上用__shfl_down的 HIP 版本,结果因为 wavefront 是 64,归约的步数和 NVIDIA 不一样,数值就对不上了。后来改成显式指定归约步数,并且用__builtin_amdgcn_...这类 intrinsic 来保证顺序,才搞定。
3.3 vime 框架的接入点
vime 是 RL 训练框架,它的职责是调度 rollout 和 training 两个阶段。要让 bitwise 一致生效,vime 需要在两个地方做改造:
第一,rollout 阶段的 logprob 计算必须走 RL-Kernel。推理引擎默认用自己的 kernel 算 logprob,vime 需要把这个计算替换成 RL-Kernel 的调用。这涉及到和推理引擎的接口对接,通常是在采样完成后,把 token 序列和 hidden states 传给 RL-Kernel 重算一遍 logprob。
第二,training 阶段的 logprob 计算也必须走 RL-Kernel。训练框架默认用 PyTorch 的log_softmax,vime 需要把这个替换成 RL-Kernel 的实现。这里要注意,训练阶段有反向传播,RL-Kernel 需要同时提供前向和反向的实现,且反向的数值行为也要和 rollout 一致(虽然 rollout 不需要反向,但为了验证一致性,前向必须一致)。
vime 的接入点通常在 policy 的 forward 函数里。我的做法是写一个 wrapper,把原来的log_softmax调用替换成rl_kernel.log_softmax,然后在配置里加一个开关,方便对比开启前后的差异。
3.4 AMD ROCm 的适配细节
ROCm 这边有几个必须注意的点:
第一,版本匹配。ROCm 的版本和 PyTorch 的版本有严格的对应关系。比如 PyTorch 2.1 通常配 ROCm 5.6 或 5.7,PyTorch 2.2 配 ROCm 5.7 或 6.0。版本不匹配会导致 kernel 编译失败或者数值行为异常。我建议用官方推荐的组合,不要自己乱配。
第二,编译选项。hipcc 的-O3和-ffast-math会影响浮点行为。为了 bitwise 一致,我建议在关键 kernel 上禁用 fast-math,用-fno-fast-math或者更细粒度的-ffp-contract=off来控制 FMA(fused multiply-add)的生成。FMA 会把a*b+c合成一条指令,精度更高但和分开算的结果不同。如果 rollout 用了 FMA 而 training 没用,bitwise 就不一致。
第三,内存对齐。AMD 的 GPU 对内存对齐比较敏感,不对齐的访问可能导致 kernel 走不同的代码路径,进而影响数值。RL-Kernel 在分配显存时要注意对齐到 256 字节或更高。
第四,MIOpen 的算法选择。MIOpen 是 AMD 的深度学习库,它有很多算法可选(比如 GEMM 的不同实现)。不同算法的数值行为可能不同。为了 bitwise 一致,需要固定 MIOpen 的算法选择,可以通过环境变量MIOPEN_FIND_MODE=normal和MIOPEN_DEBUG_FIND_ONLY_SOLVER来锁定。
4. 实操过程:从环境搭建到 bitwise 验证
4.1 环境准备与版本锁定
先把环境搭起来。我用的是 AMD MI210,ROCm 6.0,PyTorch 2.2。这套组合在我这边跑下来比较稳。
# 检查 ROCm 版本 rocm-smi # 输出里能看到 ROCm 版本和 GPU 型号 # 检查 PyTorch 是否识别到 GPU python -c "import torch; print(torch.__version__); print(torch.cuda.is_available()); print(torch.cuda.get_device_name(0))" # 注意:PyTorch 在 ROCm 上仍然用 cuda 这个命名空间,这是历史遗留版本锁定很关键。我建议用 conda 或者 docker 来固定环境,避免系统级的库版本冲突。Docker 的话,用rocm/pytorch:rocm6.0_ubuntu22.04_py3.10_pytorch_2.2这个官方镜像比较省事。
注意:如果你用的是消费级 Radeon 卡(比如 RX 7900),ROCm 的支持不如 MI 系列完整,可能需要额外的环境变量来启用。比如
HSA_OVERRIDE_GFX_VERSION=11.0.0这种,具体值取决于你的卡。这个操作有风险,可能导致不稳定,建议先在测试环境验证。
4.2 RL-Kernel 的编译与集成
RL-Kernel 通常是一个独立的库,需要单独编译。编译流程大致是:
# 克隆 RL-Kernel 仓库(假设你已经有了) cd rl-kernel mkdir build && cd build cmake .. -DCMAKE_BUILD_TYPE=Release -DGPU_TARGETS=gfx90a # gfx90a 是 MI210 的架构代号,其他卡要改 make -j$(nproc)GPU_TARGETS这个参数很重要,它决定了编译出来的 kernel 针对哪个架构优化。填错了会导致 kernel 跑不起来或者性能极差。常见的架构代号:
| GPU 型号 | 架构代号 |
|---|---|
| MI210 | gfx90a |
| MI250 | gfx90a |
| MI300 | gfx942 |
| RX 7900 XTX | gfx1100 |
| RX 6900 XT | gfx1030 |
编译完成后,把生成的.so文件放到 Python 的 path 里,或者用pip install -e .安装 Python 绑定。
集成到 vime 的时候,我写了一个简单的 wrapper:
import rl_kernel class BitwiseLogSoftmax(torch.autograd.Function): @staticmethod def forward(ctx, logits, dim=-1): # 调用 RL-Kernel 的 log_softmax output = rl_kernel.log_softmax(logits, dim) ctx.save_for_backward(output) ctx.dim = dim return output @staticmethod def backward(ctx, grad_output): output, = ctx.saved_tensors # RL-Kernel 的反向实现 grad_input = rl_kernel.log_softmax_backward(grad_output, output, ctx.dim) return grad_input, None然后在 vime 的 policy forward 里,把原来的F.log_softmax替换成BitwiseLogSoftmax.apply。
4.3 关键参数的计算与选择
这里讲几个必须算清楚的参数。
第一,归约的 block size。假设你要对长度为 4096 的序列做 sum 归约,用多大的 block?我的经验是:block size 取 256,每个线程处理 16 个元素,然后做 block-level reduction。这个配置在 AMD 上比较平衡,既不会因为 block 太小导致归约步数太多,也不会因为 block 太大导致 occupancy 下降。
计算过程:4096 / 256 = 16,每个线程处理 16 个元素。线程内的归约用顺序累加(保证顺序固定),block 内的归约用树形(log2(256) = 8 步)。这个归约树是固定的,不管输入怎么变,顺序都一样。
第二,fp32 累加的精度。fp32 有 24 位尾数,能精确表示约 1600 万个整数。对于 logprob 计算,累加的元素数量通常在几千到几万,fp32 累加的误差在 1e-6 量级,足够 bitwise 一致。如果序列特别长(比如 128k),可能需要用 Kahan summation 或者 double 累加,但那样性能会掉很多。
第三,FMA 的取舍。FMA 把a*b+c合成一条指令,精度比分开算高,但和分开算的结果不同。为了 bitwise 一致,要么两边都用 FMA,要么两边都不用。我选择在关键路径上禁用 FMA,用-ffp-contract=off编译,这样 rollout 和 training 的行为一致。
4.4 Bitwise 验证的实操流程
验证是最后一步,也是最关键的一步。我的验证流程:
准备固定输入。用固定的随机种子生成一批 logits,保存成文件。这样 rollout 和 training 用的是完全相同的输入。
分别跑 rollout 和 training 的 logprob 计算。rollout 走推理引擎 + RL-Kernel,training 走训练框架 + RL-Kernel。
对比输出的二进制。不要用
torch.allclose,那个有容差。要用torch.equal或者直接对比字节。
import torch # 加载两个阶段的输出 logprob_rollout = torch.load('logprob_rollout.pt') logprob_training = torch.load('logprob_training.pt') # bitwise 对比 if torch.equal(logprob_rollout, logprob_training): print("Bitwise 一致,通过") else: diff = (logprob_rollout - logprob_training).abs() print(f"最大差异: {diff.max().item()}") print(f"差异位置: {diff.argmax().item()}") # 找到第一个不一致的位置,方便排查 mismatch = (logprob_rollout != logprob_training).nonzero() print(f"不一致元素数量: {len(mismatch)}")- 如果不一致,逐层排查。先对比 logits,再对比 log_softmax 的中间结果(max、sum),最后定位到具体的 kernel。
我实测下来,第一次跑通常是不一致的,差异在 1e-7 到 1e-6 之间。排查下来最常见的原因是 FMA 的生成不一致,或者归约的 block size 在两边不同。把这两个固定住之后,基本就能做到 bitwise 一致。
5. 常见问题与排查技巧实录
5.1 问题速查表
| 现象 | 可能原因 | 排查方法 | 解决方案 |
|---|---|---|---|
| logprob 差异 1e-6 量级 | FMA 生成不一致 | 对比编译选项,检查-ffp-contract | 两边统一禁用或启用 FMA |
| logprob 差异 1e-3 量级 | 精度不一致(fp16 vs fp32) | 检查两边的 dtype | 关键路径统一用 fp32 累加 |
| 差异随序列长度增大 | 归约顺序不一致 | 对比 block size 和归约树 | 固定归约树,统一 block size |
| 某些 batch 一致某些不一致 | 内存对齐问题 | 检查输入 tensor 的 stride | 强制对齐到 256 字节 |
| kernel 编译失败 | 架构代号填错 | 检查GPU_TARGETS | 用rocminfo查正确的 gfx 代号 |
| 性能极差 | block size 太小或太大 | 用 rocprof 分析 | 调整到 256 或 512 |
| MIOpen 算法不稳定 | 算法自动选择 | 设置MIOPEN_FIND_MODE | 锁定 solver |
5.2 独家避坑技巧
技巧一:用torch.use_deterministic_algorithms(True)做初步筛查。这个选项会让 PyTorch 在遇到不确定的算子时报错,能帮你快速定位哪些算子有问题。但注意,它不能保证 bitwise 一致,只能保证可复现。
技巧二:保存中间结果做对比。不要只对比最终的 logprob,把 max、sum、exp 这些中间结果都存下来。这样一旦不一致,能快速定位到是哪一步出的问题。我通常会在 kernel 里加一个 debug 模式,把中间结果 dump 到文件。
技巧三:用小的输入做快速验证。不要一上来就用完整的 batch 和长序列,先用 batch=1、seq_len=128 的小输入验证。小输入下归约步数少,问题更容易暴露。验证通过后再逐步放大。
技巧四:注意 ROCm 的HSA_ENABLE_SDMA环境变量。这个变量控制是否用 SDMA 引擎做数据传输,不同设置下数据传输的顺序可能不同,进而影响数值。为了 bitwise 一致,建议固定这个变量的值。
技巧五:记录每次实验的完整环境。包括 ROCm 版本、PyTorch 版本、编译选项、环境变量。bitwise 一致很容易被环境变化打破,记录清楚才能快速复现。
5.3 一个真实的排查案例
有一次我遇到一个诡异的问题:同样的代码,在 MI210 上 bitwise 一致,在 MI250 上就不一致。排查了半天,发现是 MI250 是双 die 设计,跨 die 的通信走的是不同的路径,归约的时候如果跨了 die,顺序就和单 die 不同。
解决方案是在 kernel 里显式指定 die 的亲和性,让归约在单个 die 内完成。这个操作比较底层,需要用hipSetDevice和hipDeviceSetLimit这类 API。具体代码就不贴了,思路是:把数据先按 die 分组,每组内部做归约,最后再合并。这样归约树就固定了,不管单 die 还是双 die,结果都一样。
这个坑让我意识到,AMD 的硬件架构差异比想象中大。MI 系列和 Radeon 系列的差异、单 die 和双 die 的差异,都会影响数值行为。做 bitwise 一致的时候,必须把这些因素都考虑进去。
6. 性能与一致性的平衡:我的取舍经验
6.1 一致性的代价
做到 bitwise 一致不是没有代价的。最直接的代价是性能。
第一,禁用 FMA 会让矩阵乘的性能下降 10% 到 20%。FMA 是 GPU 上最常用的指令之一,禁用之后计算吞吐会明显下降。
第二,固定归约树会限制 kernel 的优化空间。本来可以根据输入大小动态调整 block size,现在必须固定,小输入的时候可能浪费线程,大输入的时候可能不够用。
第三,fp32 累加比 fp16 累加慢。fp16 的累加吞吐是 fp32 的两倍,但为了精度只能用 fp32。
我实测的数据:在一个 7B 模型的 RL 训练里,开启 bitwise 一致后,rollout 的吞吐下降了约 15%,training 的吞吐下降了约 8%。整体训练时间增加了约 10%。
6.2 哪些地方可以放松
不是所有地方都需要 bitwise 一致。我的经验是:只对影响 logprob 的关键路径严格要求,其他部分可以放松。
具体来说:
- log_softmax 的 max 和 sum 归约:必须 bitwise 一致。
- attention 的加权平均:如果 attention 的输出会影响 logits,那也必须一致。但如果 attention 用的是 flash attention 这类近似算法,可能需要换成精确实现。
- layer norm 的均值方差:影响 hidden states,间接影响 logits,建议一致。
- embedding lookup:只是查表,不涉及浮点运算,天然一致。
- 激活函数:exp、silu 这些,如果实现不同会有差异,建议统一。
放松的地方:
- KV cache 的存储布局:不影响数值,可以不同。
- 数据传输的顺序:只要最终结果一致,中间怎么传无所谓。
- 非关键路径的算子:比如一些辅助计算的 kernel。
6.3 一个折中方案
如果性能实在受不了,可以考虑混合方案:rollout 阶段用快速 kernel,training 阶段用精确 kernel,但在两者之间加一个校正步骤。
具体做法:rollout 算出的 logprob 存下来,training 阶段重算 logprob 后,用 rollout 的值做一次校正。校正的方式可以是直接替换,或者用某种插值。这样 training 的梯度是基于 rollout 的 logprob 算的,避免了不一致带来的偏差。
这个方案的缺点是:校正本身有开销,而且如果 rollout 的 logprob 本身有误差,校正后还是有误差。所以它只是权宜之计,不是根本解决方案。
7. 后续可以扩展的方向
这套东西跑通之后,我发现还有几个方向可以继续挖。
第一,扩展到更多算子。目前只搞定了 log_softmax 和相关的归约,attention、layer norm 这些还没完全覆盖。如果能把整个 transformer 的前向都做到 bitwise 一致,那 RL 训练的稳定性会有质的提升。
第二,自动化验证工具。现在验证还是手工的,每次改代码都要重新跑一遍对比。可以做一个 CI 流程,每次提交代码自动跑 bitwise 验证,不一致就报错。
第三,跨平台一致性。现在只保证了 AMD 平台内部的一致,如果 rollout 在 AMD 上、training 在别的平台上,能不能也做到一致?这个难度更大,因为不同平台的浮点行为差异更大。但理论上,只要归约树和精度都固定,是可以做到的。
第四,性能优化。现在为了 bitwise 一致牺牲了 10% 的性能,能不能通过更好的 kernel 设计把这个损失补回来?比如用更高效的归约算法,或者用 AMD 特有的指令集(比如 MFMA)来加速。
我个人在实际操作中的体会是:bitwise 一致这件事,说起来简单,做起来全是细节。每一个归约、每一次精度转换、每一个编译选项,都可能成为不一致的来源。但只要把关键路径固定住,把验证流程建起来,这件事是可以做到的。最重要的是,做完之后训练稳定性确实有肉眼可见的提升,之前那些莫名其妙的 loss 尖刺基本消失了。这个收益,值得那 10% 的性能代价。