大模型RL训练推理一致性:从logprob到采样参数的对齐实践
2026/9/9 4:32:45 网站建设 项目流程

做RL训练最怕什么?不是loss不降,也不是显存溢出,而是你在训练日志里看到reward一路飙升,高高兴兴把模型推上线,结果离线评测和线上表现双双扑街。更诡异的是,同一个prompt,在训练时采样出来的回答看起来还挺正常,一换到推理服务里,风格、长度、格式全变了,甚至出现大量空白和重复。这种“训练时天下无敌,推理时有心无力”的鬼故事,十有八九是训练-推理一致性(train-inference consistency)没做好。

LLM的RL(强化学习)训练本质上是在调整模型自身的生成分布,而RL所用的策略分布来自训练时的采样过程。如果推理阶段没有复现训练阶段的采样约定、概率计算方式、长度处理策略,那你在RL里辛辛苦苦学到的策略就会被“翻译走样”。这篇文章我会从“差异到底藏在哪儿”开始,结合PPO、GRPO这类常见RL框架,把logprob计算、温度/top_p、长度归一化、优势估计等关键环节逐项拆开,再给出可以落地的对齐改造方法和排查清单。适合正在做LLM后训练、RLHF/RLVR、或者天天被“训练推理不一致”折磨的算法工程师和训练平台同学参考。

1. RL训练为什么要刻意“复刻”推理行为

1.1 训练阶段本身就是一个“模拟推理”的过程

先想清楚LLM RL的本质。无论是PPO还是GRPO,训练时我们都需要从当前策略模型(policy model)里采样一组回答,然后对这些回答计算奖励,再用奖励去更新策略。也就是说,模型在训练时已经在生成文本了。这个生成动作,应该是未来部署到线上后同一模型生成动作的精确复刻。

可惜很多RL训练框架为了吞吐和稳定,默认的生成配置与线上推理配置并不一致。比如训练时为了跑得快,关闭了top_p随机采样,固定用greedy decoding;线上却开着top_p=0.9、temperature=0.7,这时候两者的输出分布根本不是一回事。更碍事的还有logprob的计算口径,线上推理服务一般直接调model.generate,压根不关心每条token的对数概率,而RL恰恰靠logprob来计算策略比率和重要性权重。一旦logprob计算得不对,整个目标函数就是错的,训练出来的策略和推理行为自然对不上。

1.2 一个让我记忆犹新的翻车案例

我接手过一个对话模型的RL训练项目,训练日志里胜率指标涨得很漂亮,但业务方反馈线上回复“变笨了”。排查了很久,最后发现三个一致性问题叠在一起:

  • 训练时用了temperature=1.0,但线上服务默认temperature=0.8,导致线上分布比训练分布更加尖锐,模型倾向的token更聚集,长尾能力大幅衰减。
  • 训练走的是batch generation,使用padding实现变长batch,但attention mask没有参与logprob的计算和过滤,导致<pad>token也被计入loss,模型学会在完整回答后继续输出<pad>
  • 奖励模型对回答长度做了soft惩罚,训练时实际生效,但线上生成时长度上限配错了,导致长回答被静默截断。

这三个问题叠加,最终效果就是“训练时reward高、线上对话质量差”。我把这些问题逐一修复后,同一条评测集上的线上一致性从原来的约78%提升到97%。所以训练-推理一致性不是一个“锦上添花”的指标,它是RL训练能落地的基本前提。

1.3 一致性差会带来哪些具体的“症状”

  • 训练指标虚高:RL训练时reward升高,但同一checkpoint换到推理服务评测,指标反而下降。
  • 采样分布漂移:训练采样多用greedy或高温度,线上用低温度,导致模型实际输出与训练优化过的输出风格不一致。
  • 文本结构和格式崩塌:在代码生成、JSON输出等任务中,训练时约束了格式,推理时忘了加同样的constraints,导致格式错误率暴增。
  • 长序列退化:训练时长度归一化策略和推理时的max_new_tokens不一致,导致模型在推理时过早停止或无法停止。
  • 概率异常:logprob计算错误造成策略比率失真,模型被错误地推高某些低频token概率,进而影响鲁棒性。

下面逐个拆解,到底哪些环节会造成这些症状。

2. 不一致藏在哪儿:逐个环节拆开看

2.1 logprob是我们和推理服务之间最隐秘的分歧点

RL训练,尤其是PPO这种actor-critic方法,需要计算每条token在当前策略下的对数概率。HuggingFace Transformers的model(input_ids, attention_mask, labels)返回的loss,或者logits,与你在推理时拿到logprob的方式并不一定等价。常见的不一致有:

  • 是否包含最后一个token的logprob:RL计算某个token被选择的logprob,通常指给定前缀后预测该token的logprob。对于一个长度为T的sequence,需要计算T个logprob,对应每个位置的预测概率。如果你比较粗心地用了labels直接shift后计算交叉熵,得到的loss是T或T-1个token的均值,但RL需要的往往是sequence_logprob = sum(logprob_1...logprob_T),其中logprob_1是预测第一个真实token的概率。
  • 是否加log_softmax:很多框架返回logits,你需要手动log_softmax(-1)才能拿到概率。小细节,但一旦在分发到多卡时忘记broadcast,全部沉默出错。
  • padding位置的logprob是否被截断或置零:训练时batch内序列长度不一致,我们通常用left padding(左侧填充)来生成,这样最后一个token在固定位置。但是在计算logprob时,必须通过attention_mask把padding部分的logprob过滤掉,或直接将padding位置的logprob设为0。否则计算总概率时会把padding token也乘进去,概率值严重失真。

实操中我最推荐的方案是统一以generate接口的底层函数为基准,写一个独立的compute_logprobs函数,显式传入input_idsattention_mask,返回每个序列的token级logprob列表。同时确保训练和推理共用同一个tokenizer和同一个log_softmax实现,避免fp16精度问题导致训练和推理的logprob对不上。

2.2 解码参数:temperature、top_p、top_k对RL是致命的

RL训练采样的目标是从当前策略中获得多样化样本,因此训练方常常把temperature设成1.0、top_p设为1.0,甚至直接用multinomial采样。但推理方通常为了稳定输出,会设置temperature=0.7、top_p=0.9。两边不一致的直接后果是:你训练的模型是在某个采样分布下被优化的,而线上使用的却是另一个更集中或者更发散的概率分布,策略自然就“错位”了。

举个例子:你训练时用top_p=1.0,模型学会了多样化探索,能够给出许多有创意的答案。线上推理时却把top_p降到0.5,模型被强制截断到少量高概率token,那些RL奖励教会它的“冒险”完全没法发挥,最后还是输出平庸内容。反过来,训练时用greedy解码,线上用高温采样,模型会退化成一个“概率分布都扭曲”的复读机。

那怎么定?我的建议是:RL训练开始前,先明确线上服务最终会使用的解码参数,让训练采样器和线上保持一致。比如线上是temperature=0.8, top_p=0.9, top_k=50,训练采样阶段就应该用同样的参数。如果你的RL确实需要提高探索多样性,可以在训练开始时暂时用更高的温度,但必须在课程学习的后期逐步退火到线上参数。否则你优化的是一个“过渡分布”,而不是“部署分布”。另外,repetition_penaltyfrequency_penalty这类参数也要参与一致性核对,因为它们会改变token的概率分布。

2.3 attention mask与padding策略:一个让所有人头痛的恶魔

RL训练里为了吞吐量,大家几乎都会用变长batch。这带来一个老问题——padding。常见的做法是右侧padding(right padding),但RL用这种方案有一个坑:如果batch内长度各不相同,right padding会导致序列末端的token位置不齐,你处理next_token_logprob时特别容易错位,会把padding token的logprob误当成真实token来算。

正确的方案是统一采用左侧padding(left padding),尤其当你的推理需要保证“最后一个token是真实token”时,left padding能够避免末尾填充。同时,无论是生成阶段还是logprob计算阶段,都要用attention_mask把padding位置的token在损失里mask掉,同时注意在计算reward时也不要让padding位置参与。

我习惯的做法是:

  • 在生成前,把prompt按固定方向padding,并记录original_lengths
  • 生成结束后,基于attention_mask裁剪出纯生成的token部分。
  • 在计算logprob时,只对“真实生成token”的位置计算,padding位置置为-inf或直接过滤。

如果不这样做,你训练时的loss大概率被padding token污染,而线上推理是没有padding的(一般单条推理),两边行为完全不同,一致性无从谈起。

2.4 长度归一化、长度惩罚和停止条件

很多RL框架会给长回答加分或者减分。如果训练与推理对“长”的定义不一致,比如训练时计算的是生成部分的token数,推理时却按照prompt+answer全长截断,那长回答的奖励就会被错误估计。另一个典型问题是长度归一化:在PPO训练里,群体相对奖励(如GRPO)有的会对序列长度做归一化,或是用长度奖励作为正则项,而线上服务可能用max_new_tokens来硬截断。若训练时允许最长回答为1024,线上max_new_tokens设为512,那么RL学出来的“适度长回答”根本不会出现,直接被截断了。

同时注意停止条件。训练时你可能希望模型输出一个包含多个字段的序列,用stop_str控制不要输出额外内容;但推理时却忘了设置相同的stop字符串,导致模型继续生成无关标记。这种在代码生成任务中特别致命——模型已经生成完代码块,训练时配合stop token停止,推理时却还在生成```或者解释性文本,最终解析失败。

2.5 奖励模型与KL散度的分布陷阱

RL训练时,通常有一个reward model来打分。但reward model打分时往往也有自己的“推理”配置,比如它也经过一个温度缩放,或者也会对长度做处理。如果reward model的训练和推理行为不一致,就相当于奖励信号本身就是漂移的。另外,PPO/GPRO里通常会对新策略和参考策略计算KL散度作为正则,避免策略偏移太远。KL散度的计算依赖参考模型(reference model)的logprob。如果参考模型和策略模型的logprob口径不一致,KL散度就是废的。这里参考模型的logprob计算同样需要复用训练阶段的compute_logprobs逻辑,不能从某种推理服务API里拿一个文本概率之类的值来用。

还有一个常见的坑:reward model也是用一个LM来做的,它返回一个标量reward。但很多人图省事,直接让reward model输出logits,再对答案末尾做一个线性层,于是reward数值就受生成长度影响很大。训练时你的模型学到的是“越长reward越高”的假规律,线上生成的回答自然开始废话连篇。出现这种情况时,最好回退到“用一个sequence-level回归头预测reward”的方式,而不是token-level的隐状态均值。

3. 实操:把训练和推理拉到同一条道上

3.1 第一步:锁定一个“唯一采样协议”

不论你用什么样的RL框架,第一件事是定义一份“采样协议”文档,里面写明训练和推理共用的参数组合。我建议包括:

  • temperature
  • top_p
  • top_k
  • repetition_penalty
  • max_new_tokens
  • min_new_tokens
  • stop_strings
  • left_padding
  • truncation_side
  • use_cache
  • dtype

然后把这一份配置同时用于训练采样和推理服务。注意不能只在训练代码里写死,推理服务也要用同一个配置文件加载。我们团队会维护一个sampling_config.yaml,训练和推理服务启动时都会读取这个文件,谁改了都要通知对方,避免“训练侧觉得无所谓改了温度,推理侧不知情”这种低级事故。

3.2 第二步:统一logprob实现,并加上单元测试

核心逻辑最好写成一个独立的函数,不要散落在训练脚本里。我用的是这样一个骨架:

import torch import torch.nn.functional as F from transformers import AutoModelForCausalLM, AutoTokenizer @torch.no_grad() def compute_logprobs( model, tokenizer, input_ids, attention_mask, seq_lens, # list of generated sequence lengths (excluding padding) ): logits = model(input_ids=input_ids, attention_mask=attention_mask).logits log_probs = F.log_softmax(logits.float(), dim=-1) # 对每个序列,根据输入token id取预测概率 # 我们通常让 logits 在位置 i 预测 token i+1 shift_log_probs = log_probs[:, :-1, :].contiguous() shift_labels = input_ids[:, 1:].contiguous() shift_attention = attention_mask[:, 1:].contiguous() batch_logprobs = torch.gather( shift_log_probs, dim=-1, index=shift_labels.unsqueeze(-1), ).squeeze(-1) # 用attention mask过滤padding token,padding位置的logprob置负无穷,后续也不参与和 batch_logprobs = batch_logprobs.masked_fill( shift_attention == 0, -float("inf"), ) # 按实际序列长度求和(不含padding) seq_logprobs = [] for i, seq_len in enumerate(seq_lens): seq_logprobs.append( batch_logprobs[i, :seq_len].sum().item() ) return seq_logprobs

这里我特意用了float()把logits转成float32再算log_softmax,避免fp16累加误差。正式项目里我还会加一个test_compute_logprobs,随机生成几个样本,用sanity check验证sum(logprob)与model.generate分配的概率一致。一旦这个函数错了,后面所有RL更新全是错的,所以这里花多时间都值得。

3.3 第三步:生成和训练走同一条代码路径

很多人的RL脚本是这样的:训练时使用model.generate采集样本,但在更新阶段又用model(input_ids)重新算一遍logprob。这里会有隐患:generate内部可能就是纯model forward加上采样代码,但如果你用的generate配置里有一个processor或者assisted decoding之类的加速操作,最终生成的那些token在更新阶段直接forward,很可能因为cache或者位置编码问题导致logprob算不准。为了稳,我建议在训练时不用那些“黑科技加速生成”的API,而是直接用最朴素的model.generate(sampling_config),并且保证采样出的token序列确实能通过model(input_ids)计算出一致的logprob。如果一定要用加速推理服务(比如vLLM)来做RL采样,你需要保证vLLM计算logprob的方式和训练时的compute_logprobs一致。这是另外一个复杂度,需要在采样的同时返回logprob,或者至少返回采样概率。目前vLLM支持传入prompt_logprobs参数,但是生成的token logprob是否与HuggingFace一致,还要仔细校准。我个人建议在早期训练阶段还是以原生HF生成为主,等所有参数都验证好了再考虑加速推理替换。

3.4 第四步:让reward计算在生成序列而不是padding序列上

你写的奖励函数也要和推理行为保持一致。比如:

  • 以生成部分为准,而不是整个batch。
  • 如果奖励会惩罚超出长度的回答,那么推理服务也要用相同的长度上限。
  • 如果奖励函数里包含“必须出现某个关键词”,推理服务的输出也要能通过同一套正则解析。

实操上,RL训练完成一次采样后,立即对每个样本做“离线推理等价性检查”:把我们采样得到的回答丢到线上推理服务,用同样的prompt再采一次样,比较两个输出的分布是否接近(至少看期望reward差异是否在0.05以内)。如果差异大,那就是某些方面不一致,应立即回滚配置。

4. PPO/GRPO训练中那些“隐性不一致”的细节

4.1 优势估计与baseline:不要忽略token级的一致性

在PPO中,我们需要估计每个token的优势值A_t = r_t + gamma * V(s_{t+1}) - V(s_t),这里的V(s_t)是critic预测的状态价值。这里的“状态”其实就是“前缀token序列”。如果你在训练时的critic网络输入是“最后一个token的隐状态”,而推理时完全不用critic(因为部署的是policy),这本身没什么,但要注意critic的输入状态与policy的状态表示要是同一个tokenizer、同一个padding方式。很多坑出现在:训练时前缀加上了<bos>,但推理时不加,所以状态价值估计和推理分布不一致。

在GRPO(Group Relative Policy Optimization)这类无critic方法中,优势是用组内相对奖励归一的。这里也要注意长度归一化的一致性。比如某个reward函数基于平均reward除以序列长度,那么训练时的长度计算必须和线上生成的token数计算一致,不能一个用字符数,一个用token数。

4.2 KL散度计算要以参考模型的logprob为准

KL散度项通常是KL(pi_theta || pi_ref),PPO里很多实现用kl = (logprob_ref - logprob_current)的近似来做。这个近似的符号方向千万别反了。理想情况是log(pi_theta / pi_ref) = logprob_theta - logprob_ref,如果logprob_current < logprob_ref,说明新策略比参考策略概率低,KL为正,loss会抑制这种偏移。

但一个常见的坑是参考模型和策略模型如果共享某些层,或者参考模型版本不对,会导致KL计算失真。一定要保证参考模型的权重是冻结的、与策略模型使用完全相同的tokenizer和logprob计算逻辑。否则KL算出来是负的相关性,策略很容易飘到奖励黑客的轨道上。

4.3 奖励模型也要走同一套推理配置

奖励模型如果也是“生成式”的(例如直接把奖励建模成最后一个答案token的概率),那训练reward model时也要注意与策略模型相同的解码配置。但更稳健的做法是reward model不生成文本,只输出一个标量,例如对最后token的hidden state做回归。这种模型不容易受到采样参数的影响,但仍受到padding和truncation的影响。建议对reward model也采用统一left padding,并在计算时传入attention_mask,避免把padding hidden state当成真实语义。

4.4 离线评测与线上一致性的差异

RL训练中我们经常用离线评测集来选checkpoint。离线评测要使用与RL训练时完全一致的采样配置。不能训练用temperature=1.0,评测用temperature=0.5,否则选出的checkpoint很可能只是在某个温度下数值好看,部署时反而变差。我建议离线评测至少做两轮:

  • 第一轮:与训练配置完全一致的采样,得到分布内的评测结果,用来监控训练是否过拟合。
  • 第二轮:与线上配置完全一致的采样,得到部署预期结果。两者都要汇报,任何一边有大的落差都要检查一致性。

4.5 多卡训练下的随机种子与batch顺序一致性

另一个隐秘的不一致来自随机种子。训练时如果你在每个rank上用不同的seed做采样,推理时只用一个seed,那么采样到的样本分布会有细微差别,尤其是使用multinomial采样时。对于RL来说,每次更新应该确保“同一条prompt在多个采样中公平覆盖”,所以训练时通常让每个rank都使用相同的global seed,只不过每个rank处理不同的prompt。推理时也无所谓seed是否相同,但最好固定下来方便复现。

另外,数据的padding顺序也要全局统一。有的框架在dataloader里自动按长度排序,但推理服务往往按请求顺序来,这就导致同一个prompt在不同batch里的实际输入token序列不同(因为padding位置不同),注意力结果理论上相同,但浮点误差可能在fp16下会造成细微差异。如果追求严格一致,可以在推理服务也采用left padding和动态batch,并且在日志里记录原始输入长度。

5. 常见问题与排查实操手册

5.1 典型问题速查表

症状可能原因排查手段修复方式
训练reward上升,线上效果下降采样参数不一致;奖励函数与线上行为不一致对比训练采样配置与线上服务配置统一解码参数和停止条件
输出重复、空白字符变多padding token被计入loss;温度或top_p差异检查训练用logprob是否带attention_mask过滤使用left padding,统一logprob计算
生成长度不受控长度归一化不一致;min/max tokens不一致核对训练和推理的max_new_tokens两端使用同一长度政策
模型倾向输出格式错误训练时使用stopping criteria,推理时未使用检查stop_words_ids是否配置将stop_words_ids同时配置到推理服务
reward初值正常但更新后崩溃KL散度符号反了;reward model logprob口径不一致检查KL项符号;复算参考模型logprob统一logprob基准
同一prompt采样结果不稳定temperature/top_p不一致;随机种子不固定固定seed,核对外部参数确定唯一sampling config
多卡训练每卡生成分布不同seed设置不当;padding顺序不同检查每个rank的seed统一设置global seed,统一padding side
部署后模型似乎“变笨”推理服务精度与训练精度不一致对比fp16/bf16下logprob差异使用bf16,或校准精度敏感性

5.2 我踩过的一个“reference model logprob”的坑

有一次PPO训练总是发生kl散度突然升高然后reward崩掉。查了很久,发现参考模型加载时,我没有冻结它,而是跟随策略模型一起进行了梯度更新(代码里忘了requires_grad_(False))。更诡异的是,因为参考模型的输出被用来计算KL,它一旦更新,KL的“锚点”就漂移了,策略模型就拼命去追一个移动的靶子,结果双双发散。这个教训就是:参考模型必须彻底冻结,并且要定期校验它的输出是否与训练开始时一致

5.3 一个检查一致性的小工具

我一般会在训练脚本里附加一个简单的“一致性自检”函数:每训练一定步数,随机选取20条prompt,分别用训练采样配置和线上服务配置采样,各生成3遍,然后计算:

  • 生成序列的平均长度
  • ROUGE-L自相似度(用来衡量输出多样性)
  • 关键词命中率

如果两个配置下这三个指标差异超过阈值,立刻告警。这个自检能提前暴露配置漂移,而不是等到上线才发现问题。

5.4 和Karpathy提到的“LLM wiki”相关的我的理解

最近很多人看Karpathy的LLM相关wiki和笔记,里面也反复强调“训练和推理的一致性”以及“正确的logprob语义”。其实他提到过一个核心观点:LLM本质上是一个概率分布模型,你在训练时用的是什么分布,推理时就应该是什么分布。这个说法很直白,但落地时牵扯到无数细节。如果你把RL训练比作“在模拟器里开车”,那推理就是“真车上路”。模拟器里的方向盘转角、油门响应、刹车距离都必须和真车一致,否则老司机也开不好。

6. 最后分享几个实战心得

做LLM RL训练一年多,我最大的体会是不要迷信复杂的训练技巧,先把“一致性”做到位,很多指标会自动变好。有几个细节我每次开新项目都会先检查:

  1. 训练脚本里采样参数的来源必须是配置文件,而不是硬编码。我会把采样相关参数全部放在一个SamplingConfigdataclass里,训练和推理各读一次,确保没有偏差。
  2. logprob计算一定要独立成函数并写单测。用一个小batch,手动构造input_ids和attention_mask,验证sum(logprob)是否等于torch.distributions.Categorical手算的结果。这个测试能挡住绝大部分低级错误。
  3. 每次模型保存后,立刻做一次“采样一致性快照”。用固定prompt,保存模型分布下的输出分布,并和推理服务的输出做对比。如果发现不一致,先查版本号、配置、tokenizer文件是否同步。
  4. 奖励函数尽量少用绝对阈值,多设计成平滑函数的奖励。比如长度惩罚,可以做成“超出区间后线性衰减”,而不是“长度大于N直接给0”,这样即使训练和推理的采样长度有微小差异,奖励也不至于突变。
  5. 不要轻视浮点精度。实验发现,在fp16下,HuggingFace的logits在长序列时可能出现比较明显的误差,而推理服务如果用fp16或bf16也会有所不同。我的经验是在logprob计算时转换为float32,并且在奖励计算中也用float32,避免因为精度导致奖励分数抖动。

这些问题做完后,你再去看训练日志和线上表现,会发现“训练reward高但线上不行”的情况少了很多。LLM RL本来就是一个容易出幺蛾子的领域,先把训练-推理一致性打牢,后面加探索、加奖励塑形、加多轮RL才会更有底气。希望这篇文章能帮你少走点弯路。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询