LLM学习笔记-偏好对齐-DPO代码实现

 

前言

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