)
train-llm-from-scratch 之 DPO 实战用一条损失替代完整 RLHF 循环含 ORPO / KTO 变体【免费下载链接】train-llm-from-scratchA straightforward method for training your LLM, from downloading data to generating text.项目地址: https://gitcode.com/GitHub_Trending/tr/train-llm-from-scratch导读本文以开源仓库 train-llm-from-scratch 的 docs/05_dpo.md 为主线系统讲解 Direct Preference OptimizationDPO如何绕过奖励模型 RL 循环的复杂管线直接在偏好对上优化策略。你将掌握DPO 的数学目标与序列对数概率实现、ORPO 与 KTO 两种变体的差异、scripts/train_dpo.py的完整训练流程与全部配置参数、命令行运行方式以及如何读懂训练日志中的每个指标。本文结合仓库源码src/post_training/dpo.py、src/post_training/rollout.py、config/post_training_config.py等逐行拆解底层原理让文章兼具实战可复制性与源码级深度。本文承接的是本仓库后训练管线中的偏好对齐阶段先经 预训练 得到基座模型、SFT 得到sft.pt再进入本阶段的 DPO 对齐。为什么需要 DPORLHF 的捷径传统 RLHF 需要三个组件协同工作先训练一个奖励模型模拟人类偏好再运行 PPO 之类的强化学习循环包含 rollout 采样、价值函数估计、重要性采样裁剪等管线长、超参多、数值不稳定。DPO 的核心洞察是把最大化隐式奖励这个 RL 目标解析地折叠进一个简单的分类式损失。策略只在偏好对chosen / rejected上直接优化配一个冻结的 SFT 模型副本作为参考锚点reference anchor。整个流程不再需要奖励模型、不再需要 rollout 采样、不再需要价值函数——只有一个干净的损失。仓库文档用一张流程图概括了 DPO / ORPO / KTO 共享的数据流diagrams/05_dpo.pngchosen / rejected 偏好对同时送入可训练的策略policy与冻结的 SFT 参考副本reference两侧分别算出序列对数概率π_chosen / π_rejected与ref_chosen / ref_rejected喂给 DPO 损失-log σ(β·Δlogratios)最后走 AdamW 更新。在仓库实现中scripts/train_dpo.py后面还通过--loss_type开关提供了两种流行变体ORPOodds-ratio preference optimization无需参考模型把 chosen 响应的 SFT 负对数似然与一个胜率比odds-ratio偏好项合成一个目标将 SFT 与对齐合并为一个阶段KTOKahneman-Tversky optimization不要求成对数据把 chosen 视为期望、rejected 视为不期望以批内估计的参考 KL 为基线适合只有赞/踩信号thumbs-up/down而非严格成对偏好的场景。共享的原料序列对数概率sequence log-probsDPO 比较的是策略让 chosen 响应比 rejected 响应更可能的程度相对于参考模型有多大提升。因此首要任务是在两个模型下分别计算每条响应的求和对数概率。仓库把这一计算收敛为 sequence_logprobs并被 PPO / GRPO 复用def sequence_logprobs(model, sequences, response_mask, *, temperature1.0, requires_gradTrue): lp, mask compute_logprobs(model, sequences, response_mask, temperaturetemperature, requires_gradrequires_grad) m mask.to(lp.dtype) return (lp * m).sum(dim-1), m.sum(dim-1) # (summed logprob, #tokens) per sequence底层实现细节见 compute_logprobs采用teacher-forcing 重算logits[:, t]预测sequences[:, t1]所以返回的张量长度为T-1与目标位置1..T-1对齐response_mask同步左移一位只对响应completion位置的 token 求对数概率并求和prompt 位置不计入对数概率始终在 fp32 下计算logits.float()即使外层处于 bf16 autocast 中——因为 DPO/PPO/GRPO 都要做对数概率相减bf16 舍入误差在这里有害requires_gradFalse时在no_grad上下文运行正好用于参考模型/旧策略快照。训练循环中chosen 与 rejected 被拼接成一批同时过模型见 train_dpo.py 的 _compute_lossesids torch.cat([batch[chosen_ids], batch[rejected_ids]], dim0) mask torch.cat([batch[chosen_mask], batch[rejected_mask]], dim0) psum, pn _logps(policy, ids, mask, requires_gradTrue) pc, pr, ncn, nrn psum[:B], psum[B:], pn[:B], pn[B:]核心目标函数dpo_loss 及其数学含义标准 DPO 目标位于 dpo_lossβ 温度控制策略推离参考模型的力度def dpo_loss(policy_chosen_logps, policy_rejected_logps, ref_chosen_logps, ref_rejected_logps, beta0.1): pi_logratios policy_chosen_logps - policy_rejected_logps ref_logratios ref_chosen_logps - ref_rejected_logps logits pi_logratios - ref_logratios loss -F.logsigmoid(beta * logits).mean() chosen_reward beta * (policy_chosen_logps - ref_chosen_logps).detach() rejected_reward beta * (policy_rejected_logps - ref_rejected_logps).detach() return loss, chosen_reward, rejected_reward逐项解读pi_logratios log π(y_w) − log π(y_l)当前策略对 chosenwwinner与 rejectedlloser响应对数概率之差衡量策略内部更偏好谁ref_logratios log π_ref(y_w) − log π_ref(y_l)参考模型冻结的 SFT 副本同样的量logits pi_logratios − ref_logratios策略相对参考模型的额外偏好增益loss −log σ(β · logits)让logits越正越好——策略应比参考模型更强烈地偏好 chosen。当logits 0策略与参考一样不区分时损失恰为log 2 ≈ 0.693这就是文档中DPO/KTO 初始损失接近 0.693的由来chosen_reward / rejected_reward即隐式奖励β·(log π − log π_ref)用.detach()断开梯度仅作为训练日志中的诊断量不参与反向传播。关于 β值越大目标越激进策略被推离参考模型越远。仓库默认beta0.1文档特别警告 DPO 要用很小的学习率默认5e-7因为很容易过度推离参考模型导致模型退化所以要温和地训练。两个变体ORPO无参考与 KTO赞/踩信号ORPO参考无关SFT 对齐一步完成orpo_loss 不需要参考模型改用per-token 均值对数概率目标为L NLL(chosen) λ · (−log σ(log_odds_chosen − log_odds_rejected)) 其中 log_odds mean_logp − log(1 − exp(mean_logp))第一项nll −chosen_mean.mean()chosen 响应上的 SFT 负对数似然保证生成质量不塌第二项or_loss胜率比偏好项推动策略相对提高 chosen 的胜率λorpo_lambda默认1.0平衡两项。代码中的_log1mexpdpo.py用于数值稳定地计算log(1 − exp(x))x0避免浮点溢出。由于 ORPO 没有参考模型训练循环中ref直接被置为None见下文 train_dpo.py 第 80 行且其日志中的隐式奖励就是两侧的均值对数概率本身。从源码结构看ORPO 也是三者中唯一不需要在每步额外做一次参考模型前向的方法显存和算力开销最低。KTO从期望/不期望信号学习kto_loss 在成对数据上模拟非成对场景chosen 视为 desirable、rejected 视为 undesirable以批内估计的参考 KL 为基线kl torch.cat([chosen_logratio, rejected_logratio]).mean().clamp(min0).detach() chosen_losses 1.0 - torch.sigmoid(beta * (chosen_logratio - kl)) rejected_losses 1.0 - torch.sigmoid(beta * (kl - rejected_logratio)) loss (desirable_weight * chosen_losses).mean() (undesirable_weight * rejected_losses).mean()kl是该 batch 内所有样本对数比值policy vs reference的均值clamp(min0)并detach()作为参考 KL 基线chosen 的损失鼓励其对数比值高于基线rejected 的损失鼓励其低于基线desirable_weight / undesirable_weight默认均为1.0可用于处理赞/踩样本量不均。三者的隐式准确率统一由 implicit_accuracy 计算(chosen_reward rejected_reward).float().mean()即策略隐式奖励更偏好 chosen 的配对比例。训练器从 sft.pt 初始化冻结参考副本逐步对齐scripts/train_dpo.py 的完整工作流加载策略load_backbone_from_ckpt(cfg, cfg.sft_ckpt, ctx.device)从sft.pt构建 Transformer 并装载权重load_backbone_from_ckpt会自动剥离 DDP 的module.前缀、丢弃奖励/价值头等非骨干键见 utils.py构造冻结参考make_frozen_copy(policy, devicectx.device)深拷贝策略、置 eval 模式并关闭全部梯度make_frozen_copyORPO 模式下ref None跳过该步骤每步计算_compute_losses(policy, ref, batch, cfg, ctx)依据cfg.loss_type分派到 dpo / orpo / kto 三个损失策略侧在amp_autocastbf16下计算且requires_gradTrue参考侧在torch.no_grad()下计算反向与更新loss.backward()→clip_grad_norm_(cfg.grad_clip)→optimizer.step()学习率由cosine_lr提供线性 warmup 余弦退火见 optim.py周期评估每eval_steps在留出的preferences_test.jsonl上计算测试隐式准确率与 margineval_implicit_acc训练结束由主进程再评估一次并保存最终检查点。对应的关键代码policy load_backbone_from_ckpt(cfg, cfg.sft_ckpt, ctx.device) ref make_frozen_copy(policy, devicectx.device) if cfg.loss_type ! orpo else None policy ddp_wrap(policy, ctx) optimizer configure_optimizer(unwrap(policy), cfg.lr, cfg.weight_decay) ... loss, cr, rr _compute_losses(policy, ref, batch, cfg, ctx) loss.backward() torch.nn.utils.clip_grad_norm_(policy.parameters(), cfg.grad_clip) optimizer.step()优化器采用标准 GPT 配方configure_optimizerAdamWbetas(0.9, 0.95)权重衰减只作用于维度 ≥2 的矩阵参数bias / LayerNorm / embedding 等 1D 参数不衰减。偏好数据从公开数据集到训练输入DPO 的输入由 prepare_preference_data.py 从真实公开数据集构建Anthropic/hh-rlhf人类 helpful/harmless 偏好对脚本按\n\nAssistant:标记切分对话得到 (prompt, response)HuggingFaceH4/ultrafeedback_binarizedLLM 评判的偏好对取chosen/rejected最后一轮内容。产出 JSONL每行{prompt, chosen, rejected}训练集写preferences.jsonl、留出测试集写preferences_test.jsonlPYTHONPATH. HF_HOME/ephemeral/hf_cache python scripts/prepare_preference_data.py \ --source both --max_per_source 40000 --out_dir /ephemeral/datapreference_dataset.py 中的迭代器负责批处理通过 chat template 把prompt response编码为 token ids response maskchosen 与 rejected 右填充到同一长度以共享一次前向模型是因果注意力最后一个真实 token 不会关注其后的 padding填充位置也被 mask 在损失中剔除并按rows[rank::world_size]在多个 DDP rank 间分片。运行 DPO命令行与完整配置参数三种损失类型各有对应的推荐命令行见 train_dpo.py 与 docs/05_dpo.mdPYTHONPATH. python scripts/train_dpo.py --loss_type dpo --beta 0.1 PYTHONPATH. python scripts/train_dpo.py --loss_type orpo --orpo_lambda 1.0 PYTHONPATH. torchrun --standalone --nproc_per_node2 scripts/train_dpo.py最后一条演示 DDP 双卡训练。CLI 由 parse_config_with_json 统一生成DPOConfig中每个字段自动变成--field参数并额外提供--config指定阶段 JSON与--print-config打印解析后的完整配置并退出。配置解析顺序低 → 高dataclass 默认值 configs/base.json 阶段 JSON 命令行--field覆盖。完整配置见 configs/dpo.json对应 DPOConfig 的全部字段参数默认值含义sft_ckpt/ephemeral/ckpts/sft.pt策略初始权重同时用于制作冻结参考副本pref_path/ephemeral/data/preferences.jsonl偏好对训练数据out_ckpt/ephemeral/ckpts/dpo.pt输出检查点路径loss_typedpodpo|orpo|ktobeta0.1DPO/KTO 的温度控制推离参考模型的力度orpo_lambda1.0ORPO 胜率比项的权重仅loss_typeorpo生效batch_size8每步的偏好对数量每对 2 条序列过模型epochs1遍历训练数据的轮数eval_steps200每隔多少步在测试集上评估隐式准确率与 marginwarmup_steps50线性预热步数lr5e-7学习率刻意很小防止过度推离参考weight_decay0.0权重衰减grad_clip1.0梯度裁剪范数max_len768单侧序列最大长度超出截断save_every500周期性保存检查点的步数间隔另有继承自BaseModelConfig的模型与运行时字段vocab_size50304、context_length1024、n_embed1024、n_head16、n_blocks24约 400M 参数的 mid 配置可在一张 H100 上跑通2×H100 上训练时间合理以及devicecuda、amp_dtypebf16、seed1337、compileFalse、use_wandbFalse等。仓库还提供了微型 SMOKE 配置configs/smoke/dpo.jsonbatch_size4、max_len256、warmup_steps2用于 CPU/单卡快速冒烟测试。提示DPO 阶段建议以 SFT 检查点 为起点即先完成上一阶段的sft.pt产出检查点/数据等重工件默认存放在/ephemeral大容量盘上路径可通过配置覆盖。读懂训练日志每个数字的含义训练过程中每 20 步打印一行指标train_dpo.py并同步到 MetricsLogger / wandblossDPO/KTO 初始接近0.693即−log σ(0)策略与参考无差异时ORPO 起始值更高因为它额外包含了 chosen 的 NLL 项acc隐式奖励准确率即批内chosen_reward rejected_reward的比例应稳定爬到0.5以上r_chosen / r_rejected隐式奖励β·(log π − log π_ref)的批均值两者之差margin应随训练扩大——这正是策略在偏好上拉开差距的直接证据周期评估还会输出test_acc / test_margin在留出测试集上而GSM8K dev 准确率是最终的下游真实检验仓库后续 评估 阶段会用到。最终检查点保存到/ephemeral/ckpts/dpo.pt即out_ckpt采用仓库统一的检查点形状model_state_dict/optimizer_state_dictstage/cfg/step/metrics元数据见 save_stage_ckpt。下一步偏好对齐完成后可继续进入基于 RL 的路径PPO 与 GRPO它们复用本文介绍的sequence_logprobs作为共享基础设施并配合奖励模型或 GSM8K 验证器进行策略优化。【免费下载链接】train-llm-from-scratchA straightforward method for training your LLM, from downloading data to generating text.项目地址: https://gitcode.com/GitHub_Trending/tr/train-llm-from-scratch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考