前言
DPO 核心原理(一句话)
用有监督的对比学习,直接优化 πθ 对 chosen 和 rejected 的相对概率,同时用 πref 作为锚点防止跑偏。
DPO 核心公式
\[L_{DPO} = -\log \sigma \left( \beta \cdot \log \frac{\pi_\theta(y_w|x)}{\pi_{ref}(y_w|x)} - \beta \cdot \log \frac{\pi_\theta(y_l|x)}{\pi_{ref}(y_l|x)} \right)\]公式拆解:
| 符号 | 含义 |
|---|---|
| $\pi_\theta$ | 训练模型(可更新) |
| $\pi_{ref}$ | 参考模型(冻结) |
| $y_w$ | 获胜回答(chosen) |
| $y_l$ | 失败回答(rejected) |
| $\beta$ | 温度系数,控制偏好敏感度 |
| $\sigma$ | sigmoid 函数 |
核心代码实现(3 个函数)
函数 1:计算序列的 log 概率
python
import torch
import torch.nn.functional as F
def get_sequence_logps(model, input_ids, attention_mask, prompt_len):
"""
计算模型对 response 部分的 log 概率
核心逻辑:
1. 模型前向 → logits
2. 只取 response 部分 (prompt_len 之后)
3. log_softmax → log 概率
4. gather 提取对应 token 的概率
5. sum 得到序列 log 概率
Args:
model: 语言模型
input_ids: (batch, seq_len)
attention_mask: (batch, seq_len)
prompt_len: int, prompt 的 token 数量
Returns:
(batch,) 每个序列的 log 概率
"""
# 1. 前向传播
outputs = model(input_ids, attention_mask=attention_mask)
logits = outputs.logits # (batch, seq_len, vocab_size)
# 2. 只取 response 部分
# 预测第 i 个 token 用的是第 i-1 个位置的 logits
response_logits = logits[:, prompt_len-1:-1, :] # (batch, resp_len, vocab_size)
# 3. log_softmax
log_probs = F.log_softmax(response_logits, dim=-1) # (batch, resp_len, vocab_size)
# 4. 获取 response token IDs
response_ids = input_ids[:, prompt_len:] # (batch, resp_len)
# 5. gather:提取每个位置对应 token 的 log 概率
token_log_probs = torch.gather(
log_probs,
dim=-1,
index=response_ids.unsqueeze(-1)
).squeeze(-1) # (batch, resp_len)
# 6. mask 掉 padding
resp_mask = attention_mask[:, prompt_len:] # (batch, resp_len)
token_log_probs = token_log_probs * resp_mask
# 7. 求和 → 序列 log 概率
batch_logps = token_log_probs.sum(dim=-1) # (batch,)
return batch_logps
核心点总结
- 为什么用
sum而不是mean? 因为 DPO 公式中 $\log \pi(y\vert{}x)$ 是整个序列的 log 概率。 - 为什么取
[:, prompt_len-1:-1, :]? 自回归模型的 shift:预测 response 第 1 个 token 用的是 prompt 最后一个 token 的 logits。
一、DPO 损失函数与单步训练代码
函数 1:计算 DPO 损失
Python
def dpo_loss(
policy_chosen_logps, # (batch,) πθ 对 chosen 的 log 概率
policy_rejected_logps, # (batch,) πθ 对 rejected 的 log 概率
ref_chosen_logps, # (batch,) πref 对 chosen 的 log 概率
ref_rejected_logps, # (batch,) πref 对 rejected 的 log 概率
beta=0.1 # 温度系数
):
"""
DPO 损失函数
数学推导:
L = -log σ(beta * (πθ_chosen - ref_chosen - πθ_rejected + ref_rejected))
直观理解:
- 希望 πθ_chosen - ref_chosen 尽可能大(chosen 相对概率提升)
- 希望 πθ_rejected - ref_rejected 尽可能小(rejected 相对概率降低)
- 两者差值越大,loss 越小
"""
# 1. 计算 log 比值(log space 中的除法)
chosen_ratios = policy_chosen_logps - ref_chosen_logps # log(πθ/πref) for chosen
rejected_ratios = policy_rejected_logps - ref_rejected_logps # log(πθ/πref) for rejected
# 2. 隐式奖励(用于监控)
chosen_rewards = beta * chosen_ratios
rejected_rewards = beta * rejected_ratios
# 3. 计算 logits
# 正值 = 模型正确偏好 chosen
# 负值 = 模型偏好反了
logits = beta * (chosen_ratios - rejected_ratios) # (batch,)
# 4. DPO loss:-log σ(logits)
loss = -F.logsigmoid(logits).mean()
# 5. 准确率(监控用)
accuracy = (chosen_rewards > rejected_rewards).float().mean()
return loss, chosen_rewards, rejected_rewards, accuracy
公式对应关系:
\[L_{DPO} = -\log \sigma \left( \beta \cdot \left[ (\log \pi_\theta(y_w) - \log \pi_{ref}(y_w)) - (\log \pi_\theta(y_l) - \log \pi_{ref}(y_l)) \right] \right)\]函数 2:单步训练
Python
def train_step(
model, # 训练模型 (πθ)
ref_model, # 参考模型 (πref,已冻结)
batch, # {chosen_ids, rejected_ids, prompt_len, ...}
optimizer, # 优化器
beta=0.1
):
"""
单步 DPO 训练
流程:
1. 计算 πθ 对 chosen/rejected 的 log 概率
2. 计算 πref 对 chosen/rejected 的 log 概率(梯度冻结)
3. 计算 DPO loss
4. 反向传播
5. 更新参数
"""
# 1. 计算 πθ 的 log 概率
policy_chosen_logps = get_sequence_logps(
model,
batch["chosen_input_ids"],
batch["chosen_attention_mask"],
batch["prompt_len"]
)
policy_rejected_logps = get_sequence_logps(
model,
batch["rejected_input_ids"],
batch["rejected_attention_mask"],
batch["prompt_len"]
)
# 2. 计算 πref 的 log 概率(梯度冻结!)
with torch.no_grad():
ref_chosen_logps = get_sequence_logps(
ref_model,
batch["chosen_input_ids"],
batch["chosen_attention_mask"],
batch["prompt_len"]
)
ref_rejected_logps = get_sequence_logps(
ref_model,
batch["rejected_input_ids"],
batch["rejected_attention_mask"],
batch["prompt_len"]
)
# 3. 计算 DPO loss
loss, chosen_rewards, rejected_rewards, accuracy = dpo_loss(
policy_chosen_logps,
policy_rejected_logps,
ref_chosen_logps,
ref_rejected_logps,
beta
)
# 4. 反向传播
loss.backward()
# 5. 梯度裁剪 + 更新参数
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
optimizer.zero_grad()
return {
"loss": loss.item(),
"accuracy": accuracy.item(),
"chosen_reward_mean": chosen_rewards.mean().item(),
"rejected_reward_mean": rejected_rewards.mean().item(),
"reward_margin": (chosen_rewards - rejected_rewards).mean().item()
}
完整训练循环(极简版)
Python
def dpo_training_loop(
model,
ref_model,
train_dataloader,
optimizer,
beta=0.1,
num_epochs=3
):
"""简化的 DPO 训练循环"""
model.train()
ref_model.eval() # 参考模型设为 eval 模式
for epoch in range(num_epochs):
for step, batch in enumerate(train_dataloader):
# 单步训练
metrics = train_step(
model,
ref_model,
batch,
optimizer,
beta
)
# 打印进度
if step % 10 == 0:
print(f"Step {step}: loss={metrics['loss']:.4f}, acc={metrics['accuracy']:.3f}")
return model
完整可运行的最小示例
Python
import torch
import torch.nn.functional as F
from transformers import AutoModelForCausalLM, AutoTokenizer
from torch.utils.data import DataLoader
# ============ 1. 准备数据 ============
data = [
{
"prompt": "什么是AI?",
"chosen": "AI是人工智能,研究如何让机器具备智能。",
"rejected": "AI就是机器学习。"
},
# ... 更多数据
]
# ============ 2. Tokenization(简化) ============
tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen2.5-0.5B")
tokenizer.pad_token = tokenizer.eos_token
def tokenize_pair(prompt, chosen, rejected):
prompt_ids = tokenizer(prompt, return_tensors="pt")["input_ids"][0]
chosen_ids = tokenizer(prompt + chosen, return_tensors="pt", padding="max_length", max_length=128)["input_ids"][0]
rejected_ids = tokenizer(prompt + rejected, return_tensors="pt", padding="max_length", max_length=128)["input_ids"][0]
return {
"chosen_input_ids": chosen_ids.unsqueeze(0),
"rejected_input_ids": rejected_ids.unsqueeze(0),
"prompt_len": len(prompt_ids)
}
# ============ 3. 加载模型 ============
model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen2.5-0.5B")
ref_model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen2.5-0.5B")
# 冻结参考模型
for param in ref_model.parameters():
param.requires_grad = False
# ============ 4. 训练 ============
optimizer = torch.optim.AdamW(model.parameters(), lr=5e-6)
for epoch in range(3):
for item in data:
batch = tokenize_pair(item["prompt"], item["chosen"], item["rejected"])
metrics = train_step(model, ref_model, batch, optimizer, beta=0.1)
print(f"loss: {metrics['loss']:.4f}, acc: {metrics['accuracy']:.3f}")
print("训练完成!")
二、核心原理总结与对照表
核心原理总结图
Plaintext
输入: (prompt, chosen, rejected)
│
▼
┌───────────────────────┐
│ Tokenization │
│ 文字 → 数字 ID │
└───────────────────────┘
│
▼
┌───────────────────────┐
│ πθ 前向传播 │
│ → log πθ(chosen) │
│ → log πθ(rejected) │
└───────────────────────┘
│
▼
┌───────────────────────┐
│ πref 前向传播 │
│ (with torch.no_grad)│
│ → log πref(chosen) │
│ → log πref(rejected)│
└───────────────────────┘
│
▼
┌───────────────────────┐
│ DPO Loss │
│ -log σ(β * (chosen │
│ - rejected 差值)) │
└───────────────────────┘
│
▼
┌───────────────────────┐
│ 反向传播 │
│ 只更新 πθ 参数 │
│ πref 保持不变 │
└───────────────────────┘
代码 - 公式对应速查表
| 数学符号 | 代码变量 |
|---|---|
| $\log \pi_\theta(y_w\vert{}x)$ | policy_chosen_logps |
| $\log \pi_\theta(y_l\vert{}x)$ | policy_rejected_logps |
| $\log \pi_{ref}(y_w\vert{}x)$ | ref_chosen_logps |
| $\log \pi_{ref}(y_l\vert{}x)$ | ref_rejected_logps |
| $\log \frac{\pi_\theta}{\pi_{ref}}(y_w\vert{}x)$ | chosen_ratios |
| $\log \frac{\pi_\theta}{\pi_{ref}}(y_l\vert{}x)$ | rejected_ratios |
| $\beta \cdot (\text{chosen_ratios} - \text{rejected_ratios})$ | logits |
| $-\log \sigma(\text{logits})$ | loss |
三、深入理解:什么是“序列的 log 概率”?
计算序列的 log 概率是整个 DPO 实现中最容易混淆的地方。
1. 数学上:什么是“序列的 log 概率”?
1.1 语言模型在算什么?
语言模型做的事情是:给定前面的 token,预测下一个 token 的概率。例如,对于序列 ["春", "风", "拂", "面"]:
- $P(\text{“风”} \vert{} \text{“春”})$ → 看到“春”后,预测“风”的概率
- $P(\text{“拂”} \vert{} \text{“春”}, \text{“风”})$ → 看到“春风”后,预测“拂”的概率
- $P(\text{“面”} \vert{} \text{“春”}, \text{“风”}, \text{“拂”})$ → 看到“春风拂”后,预测“面”的概率
整个序列的概率 = 所有位置概率的乘积(联合概率):
\[P(\text{"春风拂面"}) = P(\text{"风"}\vert{}\text{"春"}) \times P(\text{"拂"}\vert{}\text{"春风"}) \times P(\text{"面"}\vert{}\text{"春风拂"})\]1.2 为什么要用 log 概率?
因为概率都是小于 1 的数,连乘会变得极小(如 $0.1 \times 0.2 \times 0.3 = 0.006$),容易下溢(浮点数精度不够)。解决方案是取对数,让乘法变加法:
\[\log P(\text{"春风拂面"}) = \log P(\text{"风"}\vert{}\text{"春"}) + \log P(\text{"拂"}\vert{}\text{"春风"}) + \log P(\text{"面"}\vert{}\text{"春风拂"})\]这就是“序列的 log 概率”。
2. 代码中:怎么算这个值?
Python
def get_sequence_logps(model, input_ids, attention_mask, prompt_len):
"""
计算模型对 response 部分的 log 概率
"""
# ====== 第1步:模型前向传播 ======
outputs = model(input_ids, attention_mask=attention_mask)
logits = outputs.logits # shape: (batch, seq_len, vocab_size)
# ====== 第2步:只取 response 部分的 logits ======
# 预测 response 的第1个 token,用的是 prompt 最后一个 token 的 logits
# 所以 response 对应的 logits 位置是 [prompt_len-1, -1)
response_logits = logits[:, prompt_len-1:-1, :] # shape: (batch, resp_len, vocab_size)
# ====== 第3步:计算 log_softmax ======
log_probs = F.log_softmax(response_logits, dim=-1) # shape: (batch, resp_len, vocab_size)
# ====== 第4步:获取 response 的实际 token IDs ======
response_ids = input_ids[:, prompt_len:] # shape: (batch, resp_len)
# ====== 第5步:提取对应 token 的 log 概率 ======
token_log_probs = torch.gather(
log_probs,
dim=-1,
index=response_ids.unsqueeze(-1)
).squeeze(-1) # shape: (batch, resp_len)
# ====== 第6步:mask 掉 padding ======
resp_mask = attention_mask[:, prompt_len:]
token_log_probs = token_log_probs * resp_mask
# ====== 第7步:求和 ======
batch_logps = token_log_probs.sum(dim=-1) # (batch,)
return batch_logps
3. 图解:每一步在做什么
示例数据
- $\text{prompt} = \text{“问题:”}$ → token IDs:
[101, 202, 303] - $\text{response} = \text{“答案是”}$ → token IDs:
[404, 505, 606] - $\text{input_ids} = [101, 202, 303, 404, 505, 606]$
- $\text{prompt_len} = 3$
执行流程图解
Plaintext
Step 1: 模型前向传播
input_ids: [101, 202, 303, 404, 505, 606]
↓ ↓ ↓ ↓ ↓ ↓
logits: [L0] [L1] [L2] [L3] [L4] [L5]
位置: 0 1 2 3 4 5
Step 2: 只取 response 部分 [prompt_len-1 : -1] = [2 : 5]
- 预测第3个位置("答")→ 用的是位置2(":")的 logits
- 预测第4个位置("案")→ 用的是位置3("答")的 logits
- 预测第5个位置("是")→ 用的是位置4("案")的 logits
response_logits = [L2, L3, L4] ← 对应预测"答案是"
Step 3: log_softmax
对每个位置的 logits 做 softmax 再取 log
log_probs = [
log P(词 | "问题:"), # 对应预测"答"
log P(词 | "问题:答"), # 对应预测"案"
log P(词 | "问题:答案") # 对应预测"是"
]
Step 4 & 5: 获取 response_ids 并 gather 提取
response_ids = [404, 505, 606] ("答", "案", "是")
token_log_probs = [
log P("答" | "问题:"),
log P("案" | "问题:答"),
log P("是" | "问题:答案")
]
Step 6: 求和
序列 log 概率 = log P("答" | "问题:") + log P("案" | "问题:答") + log P("是" | "问题:答案")
= log P("答案是" | "问题:")
4. 为什么是 prompt_len - 1 而不是 prompt_len?
Plaintext
位置: 0 1 2 3 4 5
token: [问] [题] [:] [答] [案] [是]
↑______________↑ ↑___________↑
prompt response
(3个token) (3个token)
- 预测位置 3(“答”)需要知道位置 0, 1, 2 的信息,因此使用位置 2 的 logits。
- 预测位置 4(“案”)需要知道位置 0, 1, 2, 3 的信息,因此使用位置 3 的 logits。
- 预测位置 5(“是”)需要知道位置 0, 1, 2, 3, 4 的信息,因此使用位置 4 的 logits。
关键点:logits[i] 预测的是 input_ids[i+1]。预测 response 的第一个 token 用的是 prompt 最后一个 token 的 logits,所以起点是 prompt_len - 1,终点是 -1。
5. gather 操作详解
torch.gather 用于从每行的 log_probs 中挑选出 response_ids 指定的那一列:
Python
# 假设词表大小 = 5
log_probs = torch.tensor([
[0.1, 0.2, 0.5, 0.1, 0.1], # 位置0:各词的 log 概率
[0.2, 0.1, 0.2, 0.4, 0.1], # 位置1:各词的 log 概率
[0.3, 0.2, 0.1, 0.3, 0.1], # 位置2:各词的 log 概率
]) # shape: (3, 5)
response_ids = torch.tensor([2, 3, 1]) # 实际的 token IDs (shape: 3)
token_log_probs = torch.gather(
log_probs,
dim=-1, # 在最后一维(词表维度)上索引
index=response_ids.unsqueeze(-1) # 变成 (3, 1)
)
# 结果:
# token_log_probs = [[0.5], [0.4], [0.2]] -> squeeze 后: [0.5, 0.4, 0.2]
6. 对比:为什么用 log 概率而不是原始概率?
| 维度 | 原始概率 | log 概率 |
|---|---|---|
| 联合概率 | 连乘:$0.1 \times 0.2 \times 0.3 = 0.006$ | 连加:$-2.3 + (-1.6) + (-1.2) = -5.1$ |
| 数值稳定性 | 极小(容易下溢) | 稳定不易下溢 |
| 除法操作 | $0.1 / 0.2 = 0.5$ | $\log(0.1) - \log(0.2) = -2.3 - (-1.6) = -0.7$(除法变减法) |
7. 核心要点速查
- 序列 log 概率:每个 token log 概率的和 = $\log(\text{联合概率})$。
- 为什么求和:log 空间中的加法等于原始空间的乘法。
- 为什么取
prompt_len - 1:预测 response 第 1 个 token 用的是 prompt 最后一个 token 的 logits。 - 为什么用 gather:从词表概率分布中,提取实际 token 对应的概率。
- 为什么用 mask:忽略 padding 位置的贡献。
- DPO 需要什么:需要的是 $\log \pi(y\vert{}x)$,即整个 response 的 log 概率。
推荐阅读:
Fine-tune Mistral-7b with Direct Preference Optimization – Maxime Labonne