ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

大模型训练中的KL散度:从信息论基础到RLHF/DPO实战

大模型训练中的KL散度:从信息论基础到RLHF/DPO实战 1. 项目概述为什么大模型绕不开KL散度如果你正在接触大语言模型、扩散模型或者任何形式的生成式AI那么“KL散度”这个词一定会高频地出现在你的视野里。它可能藏在损失函数的公式里出现在模型微调的论文中或是工程师在讨论模型“对齐”时反复提及。乍一看它是个充满数学符号的距离度量让人望而生畏。但我想说的是理解KL散度是理解现代大模型如何“学习”和“被塑造”的一把钥匙。它远不止是一个数学工具而是连接模型的理论理想状态与实际可控输出之间的核心桥梁。简单来说KL散度衡量的是两个概率分布之间的“差异”或“惊喜度”。在大模型的语境下这两个分布通常一个是“野性难驯”的原始模型输出比如一个未经微调的模型可能胡言乱语另一个是我们期望的“循规蹈矩”的理想输出比如符合人类价值观的、有用的回答。训练和微调大模型的诸多目标本质上都是在用KL散度作为“缰绳” gently地有时也不那么gentle将模型从前者拉向后者。无论是让ChatGPT学会拒绝不当请求还是让Stable Diffusion生成更符合提示词的图像背后都有KL散度在默默工作。因此无论你是研究者、工程师还是深度使用者搞懂KL散度都能让你更透彻地理解模型行为甚至在调参、诊断问题时更有章法。2. 理论基石深入理解KL散度的数学本质与直观含义要驾驭一个工具必须先理解它是什么。KL散度全称Kullback-LeLeibler散度有时也叫相对熵。它的定义式对于离散分布是这样的对于两个离散概率分布P和QKL散度 D_KL(P || Q) Σ_x P(x) log(P(x) / Q(x))。对于连续分布求和就换成积分。这个公式看起来有点冷冰冰我们来给它注入一些灵魂。你可以把P想象成“真实的”或“参考的”分布把Q想象成我们用来“近似”P的模型分布。KL散度计算的是当我们用Q来编码来自P的数据时所额外付出的“信息代价”。这个“信息代价”的单位是比特如果用log2或纳特如果用自然对数ln。2.1 核心性质与关键解读理解以下几点比死记公式更重要非对称性D_KL(P || Q) ≠ D_KL(Q || P)。这是KL散度最关键、也最容易引起误解的性质。它不是距离距离是对称的。不对称性意味着方向至关重要。前向KL (P||Q)我们最小化 D_KL(P_data || Q_model)。这相当于让模型Q去“覆盖”真实数据P的所有模式。如果Q的表达能力不足比如是一个单峰高斯分布而P是多峰的Q会倾向于去“模糊地”覆盖所有峰值可能导致生成一些不伦不类的平均样本。在机器学习中极大似然估计MLE本质上等价于最小化前向KL。反向KL (Q||P)我们最小化 D_KL(Q_model || P_data)。这相当于让模型Q在真实数据P的高概率区域“扎根”同时避免去到P的低概率区域即使那些区域Q本身可能很容易生成。这会导致Q的“模式坍塌”mode dropping——它可能只抓住P的一个主要模式而忽略其他次要模式。但在生成模型中这有时反而是我们想要的因为它能避免生成低质量的、模糊的样本。非负性D_KL(P || Q) ≥ 0且当且仅当PQ时取等。这意味着差异总是正的为我们提供了一个明确的优化目标把它降到0。与交叉熵、熵的关系展开公式D_KL(P||Q) Σ P logP - Σ P logQ H(P, Q) - H(P)。其中H(P)是P的熵自身的不确定性H(P,Q)是P和Q的交叉熵。在训练中P是固定的数据分布其熵H(P)是常数。因此最小化KL散度就等价于最小化交叉熵。这就是为什么分类任务的损失函数通常是交叉熵损失——它暗含了KL散度最小化的目标。2.2 信息论视角下的直观理解想象一下你是一个天气预报员。P是真实的天气历史概率比如北京夏天60%晴30%雨10%阴。Q是你简化后的预报模型比如你偷懒总是预报80%晴20%雨0%阴。当真实天气是“阴”时P(阴)0.1你的模型Q(阴)0。log(P(阴)/Q(阴)) log(0.1/0) → 无穷大。KL散度会对这种“完全未预料到”的事件赋予极大的惩罚。这就是“零概率问题”的根源在实践比如语言模型给未登录词赋零概率中必须用平滑等技术避免。反向KL (Q||P)则像是让你这个预报员保守一点你可以只预报“晴”和“雨”完全避开“阴”这个你拿不准的类别。即使历史上确实有10%的阴天你忽略它也不会受到来自反向KL的惩罚因为Q(阴)0在求和项中该项为0。这就是“模式丢弃”。在大模型中P可能是人类标注员表现出的回答分布优质、无害、有帮助Q是我们的大模型。我们用KL散度作为约束防止模型Q偏离这个理想的分布P太远。3. 实践核心KL散度在大模型关键场景中的应用解析理论很美但落地到十亿、千亿参数的大模型中KL散度是如何具体发挥作用的呢它主要活跃在以下几个核心战场。3.1 核心战场一指导大模型微调——从RLHF到DPO大模型预训练之后其输出分布Q可能包含大量无用、有害或不一致的文本。我们希望将其对齐到人类偏好分布P。最著名的框架就是基于人类反馈的强化学习RLHF。奖励模型训练阶段虽然不直接使用KL散度但奖励模型的学习目标排序损失隐式地在学习一个能区分好坏回答的标量函数为后续阶段提供信号。强化学习微调阶段这是KL散度的主秀场。其目标函数通常是目标 期望[奖励模型打分] - β * D_KL(π_θ || π_ref)其中π_θ待微调的策略模型我们想要优化的模型。π_ref参考模型通常是微调前的SFT模型。βKL惩罚系数一个超参数。这个公式的直观解释是我们既要最大化奖励让模型输出人类喜欢的回答又要防止新模型π_θ偏离原始参考模型π_ref太远。这个KL约束项至关重要没有它模型可能会为了骗取高奖励而“走火入魔”——比如生成一堆无意义但恰好符合奖励函数模式的字符或者完全忘记之前学到的语言能力灾难性遗忘。KL散度在这里充当了正则化器确保优化过程是稳定、保守的。实操心得系数β的选择是艺术也是科学。β太大模型畏手畏脚几乎不更新对齐效果差β太小模型容易失控输出不稳定甚至退化。通常需要在一个验证集上比如看模型在无害性、有用性上的平衡进行网格搜索。一个常见的起始点是β0.1左右。直接偏好优化DPORLHF需要训练一个独立的奖励模型过程复杂。DPO提出了一种更优雅的方式它直接利用偏好数据回答A优于回答B来优化策略模型。其推导的核心妙处在于它将奖励函数用最优策略和参考策略的KL散度表示出来从而绕过了显式的奖励模型训练。DPO的损失函数直接包含了π_θ和π_ref的KL散度项。这使得微调变得像监督学习一样简单效果却可比拟RLHF成为当前个人和小团队微调大模型的首选方法之一。3.2 核心战场二控制生成过程——从核采样到指导性生成在模型推理生成文本时我们也可以通过KL散度相关的技术来控制输出的多样性和质量。核采样Top-p Sampling虽然不直接计算KL但其思想与“避免低概率词”相关可以看作是一种对模型原始分布Q的修正使其更接近一个截断后的分布P’隐式地涉及了分布差异的控制。KL惩罚解码可以在每一步生成时给候选词的概率加上一个与KL散度相关的惩罚项例如惩罚那些会导致最终序列分布偏离某个目标分布如更平缓、更多样的词。这属于更高级的生成控制技术。3.3 核心战场三模型蒸馏与压缩将一个大模型教师模型的知识迁移到一个小模型学生模型中KL散度是标准工具之一。通常我们会让学生模型去模仿教师模型的输出分布即最小化二者在相同输入下输出概率的KL散度D_KL(P_teacher || P_student)。这比单纯用硬标签one-hot训练学生模型能保留更多的“暗知识”例如不同类别之间的相对关系从而得到性能更好的小模型。3.4 核心战场四多模态与扩散模型在扩散模型中前向过程是固定的加噪过程反向过程去噪则需要学习。训练去噪网络的一个常见视角是最小化去噪后数据分布与真实数据分布之间的KL散度。而在一些多模态对齐工作中如图文匹配KL散度也可用于对齐图像编码器和文本编码器产生的特征分布。4. 实战演练动手计算与代码实现KL散度约束光说不练假把式。我们以在微调中实现KL散度惩罚为例进行一场实战。假设我们正在用PyTorch微调一个语言模型采用类似RLHF中PPO算法的简化版思想。4.1 场景设定与数据准备我们有一个参考模型ref_model例如原始的Llama-2-7b-chat一个可训练的模型trainable_model结构与ref_model相同参数从中加载。我们从一个批次batch的提示词prompts开始让trainable_model生成回答sequences并得到每个生成词元token的对数概率log_probs。4.2 关键步骤计算KL散度KL散度在序列数据上是逐词元per-token计算然后求和的。对于单个样本的单个位置我们有ref_log_probs参考模型对实际生成的那个词元的对数概率。policy_log_probs可训练模型对实际生成的那个词元的对数概率。注意我们计算的是D_KL(policy || ref)即用可训练模型作为P参考模型作为Q这与3.1节公式中的符号π_θ || π_ref一致。根据离散分布的KL散度公式kl_div policy_prob * (log(policy_prob) - log(ref_prob))由于我们已有对数概率且policy_prob exp(policy_log_prob)可以推导出数值稳定的计算方式import torch import torch.nn.functional as F def compute_kl_penalty(policy_logps, ref_logps): 计算策略模型和参考模型之间的KL散度。 policy_logps: 策略模型对生成序列的对数概率形状 [batch_size, sequence_length] ref_logps: 参考模型对同一生成序列的对数概率形状 [batch_size, sequence_length] 返回每个序列的KL散度标量对长度求平均或求和依任务而定 # 确保输入形状一致 assert policy_logps.shape ref_logps.shape # 逐词元计算KL散度 policy * (log(policy) - log(ref)) # 因为 policy_logps log(policy), ref_logps log(ref) # 所以 kl_per_token exp(policy_logps) * (policy_logps - ref_logps) # 但直接计算exp可能数值不稳定使用以下等价形式 # kl_per_token policy_logps - ref_logps # 这是 log(policy/ref) # 然后需要乘以 policy_prob。更标准的、数值稳定的做法是使用KL散度函数或如下计算 # KL(P||Q) sum_i P(i) * (log P(i) - log Q(i)) # 方法使用log_softmax和kl_div函数如果输出是logits # 但这里我们直接有对数概率假设它们已经是归一化的log_softmax后的结果。 # 一个简单且稳定的计算方法是 kl_per_token torch.exp(policy_logps) * (policy_logps - ref_logps) # 对非有效token如padding进行掩码处理 # 假设有 attention_mask有效位置为1 # attention_mask attention_mask.float() # kl_per_token kl_per_token * attention_mask.unsqueeze(-1) # 如果policy_logps是3维 # 对序列长度维度求和得到每个样本的KL散度 kl_per_sample kl_per_token.sum(dim-1) # 假设policy_logps是2维 [batch, seq_len] # 或者求平均取决于你的损失函数设计 # kl_per_sample kl_per_token.mean(dim-1) # 返回整个批次的平均KL散度 kl_mean kl_per_sample.mean() return kl_mean # 更简洁且数值稳定的实现直接使用PyTorch的kl_div函数注意输入要求 def compute_kl_penalty_stable(logits_policy, logits_ref): 使用PyTorch的F.kl_div函数。 注意F.kl_div要求输入是log-probabilitieslog_softmax后的和probabilitiessoftmax后的并且reductionbatchmean会给出真正的KL公式均值。 但我们的场景是逐词元分类分布需要先转换。 假设logits_policy和logits_ref是模型输出的原始logits [batch, seq_len, vocab_size] batch_size, seq_len, vocab_size logits_policy.shape # 计算对数概率和概率 log_prob_policy F.log_softmax(logits_policy, dim-1) # [batch, seq_len, vocab_size] prob_ref F.softmax(logits_ref, dim-1) # [batch, seq_len, vocab_size] # 计算KL散度。F.kl_div的输入顺序是inputlog-probabilities targetprobabilities # reductionnone会给出每个位置的KL kl_per_token_vocab F.kl_div(log_prob_policy, prob_ref, reductionnone, log_targetFalse) # [batch, seq_len, vocab_size] # 对词汇表维度求和得到每个token位置的KL散度 kl_per_token kl_per_token_vocab.sum(dim-1) # [batch, seq_len] # 接下来用attention_mask掩码并求平均略去mask代码 # kl_masked kl_per_token * attention_mask # kl_sum kl_masked.sum() # non_padding_tokens attention_mask.sum() # kl_mean kl_sum / non_padding_tokens return kl_per_token # 返回未掩码的具体掩码操作在外部进行4.3 整合到损失函数中在训练循环中我们的总损失大致如下# 伪代码展示逻辑 for batch in dataloader: prompts batch[input_ids] # 1. 用可训练模型生成序列并获取其logits和对数概率 outputs_trainable trainable_model(prompts, generation_config) sequences outputs_trainable.sequences logits_trainable outputs_trainable.logits # 每个位置的原始输出 # 2. 将生成的序列再次输入参考模型获取参考模型的对数概率 # 注意这里需要将生成的序列作为输入让参考模型计算每个位置的下一个词概率 with torch.no_grad(): outputs_ref ref_model(input_idssequences[:, :-1], attention_mask...) # 通常用前n-1个token预测第n个 logits_ref outputs_ref.logits # 3. 计算奖励假设有一个奖励模型这里用伪函数代替 rewards reward_model.get_reward(sequences) # [batch_size] # 4. 计算策略优势例如使用GAE这里简化 advantages compute_advantages(rewards) # [batch_size] # 5. 计算可训练模型生成序列的对数概率仅对生成的token # 我们需要获取可训练模型在生成每个token时对该token的对数概率 log_probs_trainable gather_log_probs(logits_trainable, sequences[:, 1:]) # [batch, seq_len-1] # 6. 计算参考模型的对数概率同上 log_probs_ref gather_log_probs(logits_ref, sequences[:, 1:]) # [batch, seq_len-1] # 7. 计算KL散度惩罚 kl_div compute_kl_penalty_stable(logits_trainable[:, :-1, :], logits_ref) # 注意对齐维度 kl_div_mean apply_mask_and_mean(kl_div, attention_mask) # 应用掩码并求平均 # 8. 计算策略损失例如PPO的clip损失 ratio torch.exp(log_probs_trainable - log_probs_ref.detach()) # 重要性采样比率 surr1 ratio * advantages surr2 torch.clamp(ratio, 1 - clip_epsilon, 1 clip_epsilon) * advantages policy_loss -torch.min(surr1, surr2).mean() # 9. 组合总损失 total_loss policy_loss beta * kl_div_mean # 10. 反向传播与优化 optimizer.zero_grad() total_loss.backward() optimizer.step()关键注意事项数值稳定性直接计算exp(log_prob) * (log_prob - ref_log_prob)在log_prob很小时可能导致下溢。使用F.kl_div函数或logsumexp技巧更稳妥。掩码处理必须使用attention_mask忽略填充词元padding tokens和有时忽略提示词部分确保只计算生成部分的KL散度。梯度流参考模型ref_model的参数必须用torch.no_grad()包裹确保计算KL散度时梯度不会传播到参考模型否则会破坏其作为“锚点”的作用。β系数的动态调整有些高级实现如OpenAI的原始PPO会动态调整β如果当前批次的平均KL散度与目标值如target_kl偏差太大则自动增大或减小β以维持KL散度在期望范围内。这是一个提升训练稳定性的实用技巧。5. 避坑指南KL散度实践中的常见陷阱与调优策略在实际项目中直接套用理论公式往往会踩坑。下面是我从多次实践中总结出的关键陷阱和应对策略。5.1 陷阱一KL散度爆炸或为NaN现象训练初期损失突然变成NaN或者KL散度项的值极大。根因零概率问题参考模型ref_model对某个生成词元赋予的概率为0或log概率为-inf导致计算log(policy/ref)时出现无穷大。这在词汇表很大、生成序列较长时可能发生尤其是当可训练模型“探索”到一个参考模型认为几乎不可能的词时。数值计算下溢/上溢直接计算概率的指数可能导致数值问题。解决方案概率平滑/夹紧在计算参考模型的概率时添加一个极小的epsilon如1e-8进行平滑避免零概率。或者对logits进行夹紧clamp防止出现极端的logits值。使用稳定的KL函数优先使用像F.kl_div这样经过数值优化的库函数。检查生成质量如果频繁出现NaN检查一下初始阶段模型是否生成了完全乱码。可能需要调整生成参数如降低温度或检查模型初始化。5.2 陷阱二KL约束过强或过弱现象过强KL损失项远大于策略损失项模型几乎不更新微调后性能与参考模型无异对齐失败。过弱KL损失项可以忽略不计模型迅速偏离可能产生胡言乱语或退化输出甚至忘记基础语言能力。诊断与调优监控KL曲线在训练日志中持续记录KL散度的均值。它应该从一个初始值微调开始时两模型相同KL≈0缓慢上升然后稳定在一个平台值。这个平台值就是β和任务难度共同决定的平衡点。设定目标KL范围根据经验对于对话微调每个词元的平均KL散度kl_per_token稳定在0.1~0.5纳特之间可能是合理的。如果持续高于1说明约束可能太弱如果始终接近0说明约束太强或模型没学到东西。动态β策略如前所述实现一个简单的动态调整每N步检查当前KL均值kl_mean如果kl_mean target_kl * 1.5则β β * 1.5如果kl_mean target_kl / 1.5则β β / 1.5。target_kl可以设为0.1或0.2。5.3 陷阱三参考模型的选择与更新问题参考模型应该固定吗可以用微调中的模型快照吗最佳实践初始参考模型通常使用SFT监督微调后的模型而不是原始的预训练模型。因为SFT模型已经具备了基本的指令跟随能力在此基础上进行偏好对齐更安全、更高效。固定参考模型在单次微调运行中参考模型应始终保持固定。更新参考模型会使得“锚点”移动导致优化目标混乱容易引发训练不稳定。迭代式微调如果进行多轮RLHF或DPO常见的做法是第一轮用SFT模型作为参考微调得到模型V1第二轮可以用V1作为新的参考模型继续微调。但这需要谨慎评估每一轮后的模型质量。5.4 陷阱四KL散度与奖励的尺度不匹配现象奖励模型的输出值范围例如[-10, 10]与KL散度的值范围例如[0, 2]差异巨大导致总损失被其中一项主导。解决方案奖励归一化在每个训练批次内对奖励进行减均值、除标准差的操作使其大致服从均值为0、标准差为1的分布。这能有效稳定训练。手动缩放如果奖励绝对值普遍很大可以尝试对奖励乘以一个缩放因子如0.01使其与KL散度项量级相近。观察损失组成策略损失和KL损失应在同一数量级例如都是零点几到几之间较为理想。5.5 性能优化技巧并行计算同时运行参考模型和可训练模型进行前向传播会消耗大量显存。如果显存不足可以采用串行方式先运行可训练模型生成序列并保存然后再用参考模型对这些保存的序列进行计算。这会增加时间但减少峰值显存。缓存参考模型输出对于固定的提示词库可以预先用参考模型计算其生成分布或对数概率并缓存起来在训练时直接读取节省大量计算。但这只适用于离线强化学习或某些特定设置。使用融合算子像NVIDIA的FusedAdam优化器、PyTorch的scaled_dot_product_attention等可以加速计算。对于自定义的KL计算确保使用向量化操作避免Python循环。理解并妥善处理这些陷阱你的大模型微调之旅就会平稳很多。KL散度就像一位严格的教练用好了它能引导模型走向卓越用不好则可能让训练寸步难行或彻底失控。
返回列表