ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

大模型RL训练推理一致性:从logprob到采样参数的对齐实践

大模型RL训练推理一致性:从logprob到采样参数的对齐实践 做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_p0.9、temperature0.7这时候两者的输出分布根本不是一回事。更碍事的还有logprob的计算口径线上推理服务一般直接调model.generate压根不关心每条token的对数概率而RL恰恰靠logprob来计算策略比率和重要性权重。一旦logprob计算得不对整个目标函数就是错的训练出来的策略和推理行为自然对不上。1.2 一个让我记忆犹新的翻车案例我接手过一个对话模型的RL训练项目训练日志里胜率指标涨得很漂亮但业务方反馈线上回复“变笨了”。排查了很久最后发现三个一致性问题叠在一起训练时用了temperature1.0但线上服务默认temperature0.8导致线上分布比训练分布更加尖锐模型倾向的token更聚集长尾能力大幅衰减。训练走的是batch generation使用padding实现变长batch但attention mask没有参与logprob的计算和过滤导致padtoken也被计入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的logprobRL计算某个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_ids和attention_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采样。但推理方通常为了稳定输出会设置temperature0.7、top_p0.9。两边不一致的直接后果是你训练的模型是在某个采样分布下被优化的而线上使用的却是另一个更集中或者更发散的概率分布策略自然就“错位”了。举个例子你训练时用top_p1.0模型学会了多样化探索能够给出许多有创意的答案。线上推理时却把top_p降到0.5模型被强制截断到少量高概率token那些RL奖励教会它的“冒险”完全没法发挥最后还是输出平庸内容。反过来训练时用greedy解码线上用高温采样模型会退化成一个“概率分布都扭曲”的复读机。那怎么定我的建议是RL训练开始前先明确线上服务最终会使用的解码参数让训练采样器和线上保持一致。比如线上是temperature0.8, top_p0.9, top_k50训练采样阶段就应该用同样的参数。如果你的RL确实需要提高探索多样性可以在训练开始时暂时用更高的温度但必须在课程学习的后期逐步退火到线上参数。否则你优化的是一个“过渡分布”而不是“部署分布”。另外repetition_penalty、frequency_penalty这类参数也要参与一致性核对因为它们会改变token的概率分布。2.3 attention mask与padding策略一个让所有人头痛的恶魔RL训练里为了吞吐量大家几乎都会用变长batch。这带来一个老问题——padding。常见的做法是右侧paddingright padding但RL用这种方案有一个坑如果batch内长度各不相同right padding会导致序列末端的token位置不齐你处理next_token_logprob时特别容易错位会把padding token的logprob误当成真实token来算。正确的方案是统一采用左侧paddingleft 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数推理时却按照promptanswer全长截断那长回答的奖励就会被错误估计。另一个典型问题是长度归一化在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框架第一件事是定义一份“采样协议”文档里面写明训练和推理共用的参数组合。我建议包括temperaturetop_ptop_krepetition_penaltymax_new_tokensmin_new_tokensstop_stringsleft_paddingtruncation_sideuse_cachedtype然后把这一份配置同时用于训练采样和推理服务。注意不能只在训练代码里写死推理服务也要用同一个配置文件加载。我们团队会维护一个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_idsinput_ids, attention_maskattention_mask).logits log_probs F.log_softmax(logits.float(), dim-1) # 对每个序列根据输入token id取预测概率 # 我们通常让 logits 在位置 i 预测 token i1 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, indexshift_labels.unsqueeze(-1), ).squeeze(-1) # 用attention mask过滤padding tokenpadding位置的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_{t1}) - V(s_t)这里的V(s_t)是critic预测的状态价值。这里的“状态”其实就是“前缀token序列”。如果你在训练时的critic网络输入是“最后一个token的隐状态”而推理时完全不用critic因为部署的是policy这本身没什么但要注意critic的输入状态与policy的状态表示要是同一个tokenizer、同一个padding方式。很多坑出现在训练时前缀加上了bos但推理时不加所以状态价值估计和推理分布不一致。在GRPOGroup 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训练时完全一致的采样配置。不能训练用temperature1.0评测用temperature0.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训练一年多我最大的体会是不要迷信复杂的训练技巧先把“一致性”做到位很多指标会自动变好。有几个细节我每次开新项目都会先检查训练脚本里采样参数的来源必须是配置文件而不是硬编码。我会把采样相关参数全部放在一个SamplingConfigdataclass里训练和推理各读一次确保没有偏差。logprob计算一定要独立成函数并写单测。用一个小batch手动构造input_ids和attention_mask验证sum(logprob)是否等于torch.distributions.Categorical手算的结果。这个测试能挡住绝大部分低级错误。每次模型保存后立刻做一次“采样一致性快照”。用固定prompt保存模型分布下的输出分布并和推理服务的输出做对比。如果发现不一致先查版本号、配置、tokenizer文件是否同步。奖励函数尽量少用绝对阈值多设计成平滑函数的奖励。比如长度惩罚可以做成“超出区间后线性衰减”而不是“长度大于N直接给0”这样即使训练和推理的采样长度有微小差异奖励也不至于突变。不要轻视浮点精度。实验发现在fp16下HuggingFace的logits在长序列时可能出现比较明显的误差而推理服务如果用fp16或bf16也会有所不同。我的经验是在logprob计算时转换为float32并且在奖励计算中也用float32避免因为精度导致奖励分数抖动。这些问题做完后你再去看训练日志和线上表现会发现“训练reward高但线上不行”的情况少了很多。LLM RL本来就是一个容易出幺蛾子的领域先把训练-推理一致性打牢后面加探索、加奖励塑形、加多轮RL才会更有底气。希望这篇文章能帮你少走点弯路。
返回列表