train-llm-from-scratch 之 DPO 实战:用一条损失替代完整 RLHF 循环(含 ORPO / KTO 变体)
2026/9/15 19:59:57 网站建设 项目流程

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 Optimization(DPO)如何绕过奖励模型 + RL 循环的复杂管线,直接在偏好对上优化策略。你将掌握:DPO 的数学目标与序列对数概率实现、ORPO 与 KTO 两种变体的差异、scripts/train_dpo.py的完整训练流程与全部配置参数、命令行运行方式,以及如何读懂训练日志中的每个指标。本文结合仓库源码(src/post_training/dpo.pysrc/post_training/rollout.pyconfig/post_training_config.py等)逐行拆解底层原理,让文章兼具实战可复制性与源码级深度。

本文承接的是本仓库后训练管线中的偏好对齐阶段:先经 预训练 得到基座模型、SFT 得到sft.pt,再进入本阶段的 DPO 对齐。

为什么需要 DPO:RLHF 的"捷径"

传统 RLHF 需要三个组件协同工作:先训练一个奖励模型模拟人类偏好,再运行 PPO 之类的强化学习循环(包含 rollout 采样、价值函数估计、重要性采样裁剪等),管线长、超参多、数值不稳定。

DPO 的核心洞察是:把"最大化隐式奖励"这个 RL 目标解析地折叠进一个简单的分类式损失。策略只在偏好对(chosen / rejected)上直接优化,配一个冻结的 SFT 模型副本作为参考锚点(reference anchor)。整个流程不再需要奖励模型、不再需要 rollout 采样、不再需要价值函数——只有一个干净的损失。

仓库文档用一张流程图概括了 DPO / ORPO / KTO 共享的数据流(diagrams/05_dpo.png):chosen / rejected 偏好对同时送入可训练的策略(policy)与冻结的 SFT 参考副本(reference),两侧分别算出序列对数概率π_chosen / π_rejectedref_chosen / ref_rejected,喂给 DPO 损失-log σ(β·Δlogratios),最后走 AdamW 更新。

在仓库实现中,scripts/train_dpo.py后面还通过--loss_type开关提供了两种流行变体:

  • ORPO(odds-ratio preference optimization):无需参考模型,把 chosen 响应的 SFT 负对数似然与一个胜率比(odds-ratio)偏好项合成一个目标,将 SFT 与对齐合并为一个阶段;
  • KTO(Kahneman-Tversky optimization):不要求成对数据,把 chosen 视为"期望"、rejected 视为"不期望",以批内估计的参考 KL 为基线,适合只有赞/踩信号(thumbs-up/down)而非严格成对偏好的场景。

共享的原料:序列对数概率(sequence log-probs)

DPO 比较的是:策略让 chosen 响应比 rejected 响应"更可能"的程度,相对于参考模型有多大提升。因此首要任务是在两个模型下分别计算每条响应的"求和对数概率"。

仓库把这一计算收敛为 sequence_logprobs,并被 PPO / GRPO 复用:

def sequence_logprobs(model, sequences, response_mask, *, temperature=1.0, requires_grad=True): lp, mask = compute_logprobs(model, sequences, response_mask, temperature=temperature, requires_grad=requires_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[:, t+1],所以返回的张量长度为T-1,与目标位置1..T-1对齐;
  • response_mask同步左移一位,只对响应(completion)位置的 token 求对数概率并求和,prompt 位置不计入;
  • 对数概率始终在 fp32 下计算logits.float()),即使外层处于 bf16 autocast 中——因为 DPO/PPO/GRPO 都要做对数概率相减,bf16 舍入误差在这里有害;
  • requires_grad=False时在no_grad上下文运行,正好用于参考模型/旧策略快照。

训练循环中,chosen 与 rejected 被拼接成一批同时过模型(见 train_dpo.py 的 _compute_losses):

ids = torch.cat([batch["chosen_ids"], batch["rejected_ids"]], dim=0) mask = torch.cat([batch["chosen_mask"], batch["rejected_mask"]], dim=0) psum, pn = _logps(policy, ids, mask, requires_grad=True) 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, beta=0.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):当前策略对 chosen(w=winner)与 rejected(l=loser)响应对数概率之差,衡量策略"内部"更偏好谁;
  • 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()断开梯度,仅作为训练日志中的诊断量,不参与反向传播。

关于 β:值越大,目标越"激进",策略被推离参考模型越远。仓库默认beta=0.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)平衡两项。

代码中的_log1mexp(dpo.py)用于数值稳定地计算log(1 − exp(x))(x<0),避免浮点溢出。由于 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(min=0).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(min=0)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 的完整工作流:

  1. 加载策略load_backbone_from_ckpt(cfg, cfg.sft_ckpt, ctx.device)sft.pt构建 Transformer 并装载权重(load_backbone_from_ckpt会自动剥离 DDP 的module.前缀、丢弃奖励/价值头等非骨干键,见 utils.py);
  2. 构造冻结参考make_frozen_copy(policy, device=ctx.device)深拷贝策略、置 eval 模式并关闭全部梯度(make_frozen_copy);ORPO 模式下ref = None跳过该步骤;
  3. 每步计算_compute_losses(policy, ref, batch, cfg, ctx)依据cfg.loss_type分派到 dpo / orpo / kto 三个损失;策略侧在amp_autocast(bf16)下计算且requires_grad=True,参考侧在torch.no_grad()下计算;
  4. 反向与更新loss.backward()clip_grad_norm_(cfg.grad_clip)optimizer.step(),学习率由cosine_lr提供线性 warmup + 余弦退火(见 optim.py);
  5. 周期评估:每eval_steps在留出的preferences_test.jsonl上计算测试隐式准确率与 margin(eval_implicit_acc);训练结束由主进程再评估一次并保存最终检查点。

对应的关键代码:

policy = load_backbone_from_ckpt(cfg, cfg.sft_ckpt, ctx.device) ref = make_frozen_copy(policy, device=ctx.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_optimizer):AdamW(betas=(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_binarized:LLM 评判的偏好对,取chosen/rejected最后一轮内容。

产出 JSONL,每行{"prompt", "chosen", "rejected"},训练集写preferences.jsonl、留出测试集写preferences_test.jsonl

PYTHONPATH=. HF_HOME=/ephemeral/hf_cache python scripts/prepare_preference_data.py \ --source both --max_per_source 40000 --out_dir /ephemeral/data

preference_dataset.py 中的迭代器负责批处理:通过 chat template 把prompt + response编码为 token ids + response mask,chosen 与 rejected 右填充到同一长度以共享一次前向(模型是因果注意力,最后一个真实 token 不会关注其后的 padding,填充位置也被 mask 在损失中剔除),并按rows[rank::world_size]在多个 DDP rank 间分片。

运行 DPO:命令行与完整配置参数

三种损失类型各有对应的推荐命令行(见 train_dpo.py 与 docs/05_dpo.md):

PYTHONPATH=. 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_node=2 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_type"dpo""dpo"|"orpo"|"kto"
beta0.1DPO/KTO 的温度,控制推离参考模型的力度
orpo_lambda1.0ORPO 胜率比项的权重(仅loss_type="orpo"生效)
batch_size8每步的偏好对数量(每对 2 条序列过模型)
epochs1遍历训练数据的轮数
eval_steps200每隔多少步在测试集上评估隐式准确率与 margin
warmup_steps50线性预热步数
lr5e-7学习率(刻意很小,防止过度推离参考)
weight_decay0.0权重衰减
grad_clip1.0梯度裁剪范数
max_len768单侧序列最大长度(超出截断)
save_every500周期性保存检查点的步数间隔

另有继承自BaseModelConfig的模型与运行时字段:vocab_size=50304context_length=1024n_embed=1024n_head=16n_blocks=24(约 400M 参数的 mid 配置,可在一张 H100 上跑通,2×H100 上训练时间合理),以及device="cuda"amp_dtype="bf16"seed=1337compile=Falseuse_wandb=False等。仓库还提供了微型 SMOKE 配置(configs/smoke/dpo.jsonbatch_size=4max_len=256warmup_steps=2)用于 CPU/单卡快速冒烟测试。

提示:DPO 阶段建议以 SFT 检查点 为起点,即先完成上一阶段的sft.pt产出;检查点/数据等重工件默认存放在/ephemeral大容量盘上,路径可通过配置覆盖。

读懂训练日志:每个数字的含义

训练过程中每 20 步打印一行指标(train_dpo.py),并同步到 MetricsLogger / wandb:

  • loss:DPO/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_dict+stage/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),仅供参考

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

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

立即咨询