在语言模型对齐人类偏好的技术演进中,传统强化学习方案(RLHF)因涉及奖励模型、策略网络及价值网络的协同训练,不仅流程繁琐,还常伴随训练不稳定与算力开销大的问题。直接偏好优化(DPO)则提供了一条更为轻量化的路径:它跳过显式奖励建模,直接利用偏好数据对策略模型进行优化。本文将围绕一个极简 DPO 教学 Demo 的 Python 实现展开,基于 Hugging Face Transformers 与 Qwen 模型,带你从零理解其数据构建、损失计算与训练推理全流程。

DPO 核心原理:一个巧妙的损失函数

DPO 的核心在于其独特的损失设计,它需要两个结构相同但参数状态不同的模型:

  • 参考模型(Reference Model):通常是预训练基座模型(如 Qwen/Qwen1.5-1.8B-Chat),其参数在训练中完全冻结,作为衡量策略模型偏移量的基线。
  • 策略模型(Policy Model):与参考模型结构一致,但参数可训练,是最终用于对齐人类偏好的目标模型。

对于每一条偏好样本 (prompt, chosen, rejected),我们需要分别计算策略模型与参考模型对优选回答(chosen)和拒绝回答(rejected)的对数概率。其损失函数定义如下:

loss = -log σ( β * [ (π_logp_chosen - π_logp_rejected) - (ref_logp_chosen - ref_logp_rejected) ] )

其中 σ 为 sigmoid 函数,β 是控制策略模型偏离参考模型程度的超参数。直观理解,该损失会促使策略模型在“优选回答相对拒绝回答的偏好增益”上超越参考模型的表现,从而让模型的输出分布更贴近人类偏好。这种设计将复杂的奖励建模过程简化为一次前向传播,极大降低了训练成本。

代码结构总览与核心模块划分

整个 dpo_minimal_demo.py 脚本结构清晰,主要可拆解为以下几个部分:

  1. 环境初始化:设置环境变量、屏蔽警告、导入依赖库(如 PyTorch、Transformers)。
  2. 常量与工具函数:定义 Qwen 聊天模板、分词辅助、对数概率计算及损失函数。
  3. 数据封装:通过 PreferenceItemPreferenceDataset 组织偏好数据。
  4. 主流程:涵盖设备检测、模型加载、训练循环与推理对比。
  5. 参数解析:支持通过命令行自定义模型、轮数、学习率等超参数。

数据准备:偏好对的定义与封装

Demo 采用了一个轻量级的数据容器 PreferenceItem,用于存储单条偏好样本:用户指令(prompt)、人类更倾向的回答(chosen)以及被拒绝的回答(rejected)。

@dataclass
class PreferenceItem:
    prompt: str
    chosen: str
    rejected: str

示例数据集中硬编码了五条中文偏好对,例如:

  • prompt: “将‘你好世界’翻译成英文”
  • chosen: “Hello World”
  • rejected: “Hi Universe”

随后通过 PreferenceDataset 类将这些数据封装为标准的 PyTorch Dataset 格式,便于后续 DataLoader 迭代读取。

class PreferenceDataset(Dataset):
    def __init__(self, items: List[PreferenceItem]): ...
    def __len__(self): ...
    def __getitem__(self, idx): ...

数据加载器默认设置 batch_size=1,意味着每个训练步骤仅处理一个偏好对。这种设计便于观察梯度变化,但实际工程中建议适当增大 batch size 以提升训练效率。

loader = DataLoader(dataset, batch_size=args.batch_size, shuffle=True)

核心函数拆解:从对数概率到损失计算

5.1 计算响应序列的对数概率

函数 sequence_logprob 是 DPO 训练中最关键的计算单元。其实现逻辑可分解为以下步骤:

  1. 格式化 Prompt:利用 QWEN_CHAT_TEMPLATE 将原始指令包装成模型期望的对话格式。
  2. 拼接完整文本:将格式化后的 prompt、response 及结束符拼接为完整输入。
  3. 分词与长度对齐:分别对 prompt 和完整文本进行分词,记录 prompt 部分的 token 长度。
  4. 前向传播:将完整 token 序列输入模型,获取每个位置的 logits。
  5. 提取响应部分:从 logits 中切片出 response 对应的部分,并计算 log_softmax。
  6. 归一化处理:对响应部分的 log 概率求和后除以长度,实现长度归一化,避免长回答主导梯度。
def dpo_loss(pi_logp_chosen, pi_logp_rejected, ref_logp_chosen, ref_logp_rejected, beta=0.1):
    diff = (pi_logp_chosen - pi_logp_rejected) - (ref_logp_chosen - ref_logp_rejected)
    return -torch.nn.functional.logsigmoid(beta * diff).mean()

5.2 DPO 损失函数实现

dpo_loss 函数严格遵循 DPO 论文公式。其中 diff 表示策略模型相对参考模型的偏好增益。当 diff > 0 时,意味着策略模型对优选回答的偏好程度已超越参考模型,此时损失值较小;反之则损失增大,驱动模型向正确方向优化。

device = "cuda" if torch.cuda.is_available() else "cpu"

提示:虽然 Demo 中提供了 tokenize_batch 辅助函数(支持 padding 与 truncation),但主流程并未调用,读者可参考其实现用于批量推理场景。

训练流程与关键实现细节

训练主循环的设计体现了 DPO 的工程实践要点。首先,代码会自动检测 GPU 并启用 CUDA 加速,同时固定随机种子以保证实验可复现性。

tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
if tokenizer.pad_token is None:
    tokenizer.pad_token = tokenizer.eos_token

模型加载阶段包含两个关键操作:一是将参考模型立即切换为 eval() 模式并冻结所有参数;二是策略模型保持 train() 模式,仅更新其参数。此外,由于 Qwen 模型未定义 pad_token,代码巧妙地复用了 eos_token 作为填充符。

for epoch in range(args.epochs):
    for step, batch in enumerate(loader):
        with torch.no_grad():
            ref_logp_c = sequence_logprob(ref_model, tokenizer, prompt, chosen, device)
            ref_logp_r = sequence_logprob(ref_model, tokenizer, prompt, rejected, device)
        optim.zero_grad()
        with torch.autocast(device_type=device, dtype=dtype):
            pi_logp_c = sequence_logprob(pi_model, tokenizer, prompt, chosen, device)
            pi_logp_r = sequence_logprob(pi_model, tokenizer, prompt, rejected, device)
            loss = dpo_loss(pi_logp_c, pi_logp_r, ref_logp_c, ref_logp_r, beta=0.1)
        loss.backward()
        torch.nn.utils.clip_grad_norm_(pi_model.parameters(), 1.0)
        optim.step()

训练循环中的几个核心细节:

  • 梯度隔离:参考模型的 log 概率计算包裹在 torch.no_grad() 中,避免不必要的计算图构建。
  • 混合精度:策略模型的前向计算使用 autocast 上下文,有效降低显存占用。
  • 梯度裁剪:设置最大范数为 1.0,防止梯度爆炸导致训练不稳定。
  • 动态监控:每个 step 打印四个 log 概率值及其差值,便于实时观察模型学习动态。
for ptxt in prompts_for_inference:
    formatted_prompt = QWEN_CHAT_TEMPLATE.format(instruction=ptxt)
    inputs = tokenizer(formatted_prompt, return_tensors="pt").to(device)
    # 参考模型生成
    gen_ref = ref_model.generate(**inputs, max_new_tokens=64)
    # 策略模型生成
    gen_pi = pi_model.generate(**inputs, max_new_tokens=64)
    # 解码并打印

超参数选择上,Demo 采用学习率 1e-6 和 β 值 0.1。较小的学习率确保模型稳定更新,而 β 控制对参考模型的偏离程度——过大易导致模型遗忘原有能力,过小则对齐效果不显著。

推理对比与运行指南

训练结束后,脚本会对若干测试 prompt 分别使用参考模型与策略模型进行生成,并打印输出结果。由于 Demo 仅使用 5 条数据且训练轮次有限,你可能观察到差异不明显;但若扩充数据规模并增加训练轮数,策略模型的输出将逐渐向 chosen 回答风格靠拢。

python dpo_minimal_demo.py --model Qwen/Qwen1.5-1.8B-Chat --epochs 3 --lr 1e-6 --batch_size 1

运行脚本非常简单,支持以下参数自定义:

  • --model:模型 ID,默认 Qwen/Qwen1.5-1.8B-Chat
  • --epochs:训练轮数,默认 3
  • --lr:学习率,默认 1e-6
  • --batch_size:批次大小,默认 1
  • --max_len:最大序列长度,默认 512
import os # 导入 os 模块
import math # 导入 math 模块
import warnings # 导入 warnings 模块
from dataclasses import dataclass # 导入 dataclass 模块
from typing import List, Dict # 导入 List 和 Dict 类型
import argparse  # 将 argparse 导入到文件顶部
# 忽略 torch.cuda 模块的 FutureWarning 警告
warnings.filterwarnings("ignore", category=FutureWarning, module="torch.cuda")
# 忽略 NVIDIA GeForce RTX 警告
warnings.filterwarnings("ignore", message=".*NVIDIA GeForce RTX.*", category=UserWarning)
# 忽略 Torch 未编译 flash attention 警告
warnings.filterwarnings("ignore", message=".*Torch was not compiled with flash attention.*", category=UserWarning)
# 忽略 generation flags 警告
warnings.filterwarnings("ignore", message=".*generation flags are not valid and may be ignored:.*", category=UserWarning)
# 忽略 torch_dtype 警告
warnings.filterwarnings("ignore", message=".*`torch_dtype` is deprecated! Use `dtype` instead!.*", category=FutureWarning)
# 忽略 TF_ENABLE_ONEDNN_OPTS 警告
os.environ["TF_ENABLE_ONEDNN_OPTS"] = "0"
# 导入 torch 模块 用于张量操作
import torch
# 导入 torch.utils.data 模块 用于数据集和数据加载器
from torch.utils.data import Dataset, DataLoader
# 导入 transformers 模块 用于自动分词器和因果语言模型
from transformers import AutoTokenizer, AutoModelForCausalLM
# 导入 torch.optim 模块 用于优化器
from torch.optim import AdamW
# Qwen 聊天模板
QWEN_CHAT_TEMPLATE = (
    "<|im_start|>user\n{instruction}<|im_end|>\n"
    "<|im_start|>assistant\n"
)
# 批量编码文本
def tokenize_batch(tokenizer:AutoTokenizer, texts: List[str], max_len: int = 512) -> Dict[str, torch.Tensor]:
    """
    使用分词器批量编码文本。
    参数:
    tokenizer (AutoTokenizer): 自动分词器实例。
    texts (List[str]): 输入文本列表。
    max_len (int, 可选): 最大序列长度,默认值为 512。
    返回:
    Dict[str, torch.Tensor]: 包含编码后的张量的字典,键包括 'input_ids'、'attention_mask' 等。
    """
    out = tokenizer(texts, padding=True, truncation=True, max_length=max_len, return_tensors="pt", add_special_tokens=False)
    # padding=True 自动填充序列到 max_len
    # truncation=True 截断序列到 max_len
    # return_tensors="pt" 返回 PyTorch 张量
    # add_special_tokens=False 不添加特殊 token
    return {
        k: v for k,
        v in out.items()
    }
# 偏好数据项
@dataclass
class PreferenceItem:
    prompt: str # 存储用户输入的提示或问题
    chosen: str # 存储人类偏好的、高质量的回答
    rejected: str # 存储人类不偏好的、低质量的回答
# 内存中的偏好数据集
class PreferenceDataset(Dataset):
    # 接收并存储 PreferenceItem 实例的列表
    def __init__(self, items: List[PreferenceItem]):
        self.items = items # 存储 PreferenceItem 实例的列表
    # 返回数据集中项目的总数
    def __len__(self):
        return len(self.items) # 返回数据集中项目的总数
    # 根据索引返回单个数据项,转换为字典格式
    def __getitem__(self, idx):
        item = self.items[idx] # 获取索引为 idx 的 PreferenceItem 实例
        return {
            "prompt": item.prompt,
            "chosen": item.chosen,
            "rejected": item.rejected
        }
# 计算模型对响应的 log 概率
def sequence_logprob(model: AutoModelForCausalLM, tokenizer: AutoTokenizer, original_prompt: str, response: str, device: str) -> torch.Tensor:
    """
    计算模型对响应的 log 概率。
    参数:
    model (AutoModelForCausalLM): 自动因果语言模型实例。
    tokenizer (AutoTokenizer): 自动分词器实例。
    original_prompt (str): 用户输入的原始提示或问题。
    response (str): 模型生成的响应。
    device (str): 计算设备,如 'cuda' 或 'cpu'。
    返回:
    torch.Tensor: 响应的 log 概率张量。
    """
    # 模板应用:使用 Qwen 聊天模板格式化原始提示,确保输入格式与模型训练时一致
    formatted_prompt = QWEN_CHAT_TEMPLATE.format(instruction=original_prompt)
    # 文本构建:将格式化的提示、响应和结束 token 连接成完整文本,模拟模型生成的完整序列
    full_text = formatted_prompt + response + tokenizer.eos_token # 合并提示、响应和结束 token
    # 完整文本分词:将完整文本转换为 token ID 张量,并移动到指定设备
    full_input_ids = tokenizer(full_text, add_special_tokens=False, return_tensors="pt")["input_ids"].to(device)
    # 提示部分分词:单独对提示部分分词,计算其长度,用于后续确定响应部分的起始位置
    prompt_input_ids = tokenizer(formatted_prompt, add_special_tokens=False, return_tensors="pt")["input_ids"].to(device)
    # 响应部分分词:从完整文本中提取响应部分的 token ID 张量,用于计算 log 概率
    prompt_len = prompt_input_ids.size(1)
    # 模型前向传播:计算完整文本的 logits
    out = model(input_ids=full_input_ids)
    # 提取 logits:从模型输出中提取 logits,移除最后一个 token 的 logits
    logits = out.logits[:, :-1, :]
    # 标签构建:将完整文本的 token ID 张量向右移动一个位置,用于计算 log 概率
    labels = full_input_ids[:, 1:]
    # 计算响应部分在labels序列中的起始索引
    # 减1是因为labels已经向右移动了一位(labels = full_input_ids[:, 1:])
    start_idx = prompt_len - 1
    # 计算响应部分的长度
    response_len = labels.size(1) - start_idx
    # 检查响应部分是否为空
    if response_len <= 0:
        # 如果响应部分为空,返回一个非常小的 log 概率,避免计算错误
        return torch.tensor(-1e9, device=device)
    # 使用切片操作提取响应部分的logits
    # : 表示保留所有批次,start_idx: 表示从start_idx开始到末尾,: 表示保留所有词汇表维度
    token_logits = logits[:, start_idx:, :]
    # 使用切片操作提取响应部分的标签
    # : 表示保留所有批次,start_idx: 表示从start_idx开始到末尾,: 表示保留所有词汇表维度
    token_labels = labels[:, start_idx:]
    # 对每个 token 计算响应部分的对数概率分布
    # dim=-1 表示在最后一个维度上计算 log softmax,即对每个 token 计算其在词汇表中的概率分布
    log_probs = torch.nn.functional.log_softmax(token_logits, dim=-1)
    # 提取响应部分的 log 概率
    # 从 log 概率分布中提取对应标签的 log 概率
    # token_labels.unsqueeze(-1) 表示在最后一个维度上添加一个维度,用于匹配 gather 函数的索引要求
    # sel 表示提取的 log 概率张量,形状为 (batch_size, response_len)
    sel = torch.gather(log_probs, dim=-1, index=token_labels.unsqueeze(-1)).squeeze(-1)
    # 计算响应部分的 log 概率总和
    return sel.sum(dim=1) / response_len # 使用 response_len 归一化,更准确
# DPO 损失函数
def dpo_loss(
    pi_logp_chosen: torch.Tensor,  # 策略模型(待训练模型)对优选回答的对数概率
    pi_logp_rejected: torch.Tensor,  # 策略模型(待训练模型)对拒绝回答的对数概率
    ref_logp_chosen: torch.Tensor,  # 参考模型(如 GPT-3)对优选回答的对数概率
    ref_logp_rejected: torch.Tensor,  # 参考模型(如 GPT-3)对拒绝回答的对数概率
    beta: float = 0.1 # 超参数,控制策略模型与参考模型之间的平衡
) -> torch.Tensor:
    # (策略模型偏好差) - (参考模型偏好差)
    diff = (pi_logp_chosen - pi_logp_rejected) - (ref_logp_chosen - ref_logp_rejected)
    # 低学习率可以保证稳定性
    return -torch.nn.functional.logsigmoid(beta * diff).mean()
# 从模型输出中提取最后一个非思考标签的响应
def strip_qwen_think_tags(text: str) -> str:
    import re
    cleaned_text = re.sub(r'<\|im_start\|>thought\\n.*?<\|im_end\|>', '', text, flags=re.DOTALL) # 移除思考标签的文本
    cleaned_text = re.sub(r'(.*?<\/think>)?', '', cleaned_text, flags=re.DOTALL) # 移除思考标签
    cleaned_text = cleaned_text.replace('<|im_end|>', '') # 移除结束 token
    cleaned_text = cleaned_text.replace(QWEN_CHAT_TEMPLATE.format(instruction=''), '') # 移除聊天模板
    cleaned_text = cleaned_text.strip() # 移除首尾空格
    lines = [line.strip() for line in cleaned_text.split('\n') if line.strip()] # 移除空行
    if lines:
        return lines[-1] # 返回最后一个非思考标签的响应
    return cleaned_text
# 主函数
def main(args):
    """
    DPO 极简教学 Demo 的主函数。
    """
    # 检测并选择可用的计算设备(GPU 或 CPU)
    device = "cuda" if torch.cuda.is_available() else "cpu"
    print(f"当前运行设备: {device.upper()}")
    # 设置随机种子
    import random # 导入随机数生成模块
    import numpy as np # 导入 numpy 库,用于处理数组操作
    random.seed(args.seed) # 设置随机种子,确保实验可重复性
    np.random.seed(args.seed) # 设置 numpy 随机种子,确保实验可重复性
    torch.manual_seed(args.seed) # 设置 torch 随机种子,确保实验可重复性
    if torch.cuda.is_available(): # 如果有可用的 GPU
        torch.cuda.manual_seed_all(args.seed) # 设置所有 GPU 的随机种子,确保实验可重复性
    model_name = args.model # 从参数中获取模型名称
    tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True) # 加载模型的分词器,设置 trust_remote_code 为 True 以支持自定义代码
    if tokenizer.pad_token is None: # 如果分词器没有 pad_token
        tokenizer.pad_token = tokenizer.eos_token # 将 pad_token 设置为 eos_token
    print(f"Tokenizer 加载完成。pad_token: {tokenizer.pad_token}, eos_token: {tokenizer.eos_token}")
    # 定义偏好数据集
    items = [
        PreferenceItem(prompt="将'你好世界'翻译成英文", chosen="Hello World", rejected="Hi Universe"),
        PreferenceItem(prompt="把下面句子改礼貌:把会议纪要发我", chosen="请在方便时将会议纪要发送给我,谢谢。", rejected="把会议纪要发给我。"),
        PreferenceItem(prompt="一句话解释损失掩码", chosen="损失掩码让模型只在答案区域计算损失,专注学习输出。", rejected="损失掩码是一个很重要的东西。"),
        PreferenceItem(prompt="用一句话描述 DPO 的作用", chosen="DPO 直接优化模型以匹配人类偏好,无需奖励模型。", rejected="DPO 是一个复杂的强化学习算法。"),
        PreferenceItem(prompt="请用一句话概括 DPO 的优点", chosen="DPO 简化了对齐过程,训练稳定且计算效率高。", rejected="DPO 需要大量计算资源和复杂的奖励模型。"),
    ]
    # 初始化偏好数据集
    dataset = PreferenceDataset(items)
    # 初始化数据加载器
    loader = DataLoader(dataset, batch_size=args.batch_size, shuffle=True)
    print(f"\n加载模型: {model_name}...")
    # 加载参考模型
    # 性能优先 :优先使用更高效的 bfloat16 数据类型
    # 兼容性保障 :在不支持 bfloat16 的设备上自动回退到 float32
    dtype = torch.bfloat16 if device == "cuda" and torch.cuda.is_bf16_supported() else torch.float32
    print(f"使用的数据类型: {dtype}")
    # 加载一个预训练的因果语言模型作为参考模型
    # model_name :参考模型的名称,这里使用与策略模型相同的模型
    # trust_remote_code :设置为 True 以支持自定义代码
    # torch_dtype :指定数据类型,这里根据设备和支持情况选择 bfloat16 或 float32
    ref_model = AutoModelForCausalLM.from_pretrained(model_name, trust_remote_code=True, torch_dtype=dtype).to(device)
    ref_model.eval() # 切换参考模型为评估模式
    for p in ref_model.parameters(): # 冻结参考模型的参数,不进行训练
        # 告诉 PyTorch 不需要追踪该张量的操作,不计算梯度
        p.requires_grad_(False) # 冻结参考模型的参数,不进行训练
    print(f"参考模型已加载到 {device}")
    # 加载一个预训练的因果语言模型作为参考模型
    # model_name :参考模型的名称,这里使用与策略模型相同的模型
    # trust_remote_code :设置为 True 以支持自定义代码
    # torch_dtype :指定数据类型,这里根据设备和支持情况选择 bfloat16 或 float32
    pi_model = AutoModelForCausalLM.from_pretrained(model_name, trust_remote_code=True, torch_dtype=dtype).to(device)
    # 切换策略模型为训练模式
    pi_model.train()
    print(f"策略模型已加载到 {device}")
    # 使用 args 中的学习率初始化 AdamW 优化器
    # pi_model.parameters() :策略模型的所有可训练参数
    # lr=args.lr :学习率,从参数中获取
    optim = AdamW(pi_model.parameters(), lr=args.lr)
    print(f"\n{'='*20} 开始 DPO 训练 {'='*20}")
    # 训练循环,遍历 args.epochs 轮
    for epoch in range(args.epochs):
        total_loss = 0.0 # 累计每个 epoch 的损失
        print(f"\n{'='*20} Epoch {epoch+1}/{args.epochs} {'='*20}")
        # 遍历数据加载器中的每个批次
        for step, batch in enumerate(loader):
            item_prompt = batch["prompt"][0] # 从批次中提取提示文本
            item_chosen = batch["chosen"][0] # 从批次中提取选中文本
            item_rejected = batch["rejected"][0] # 从批次中提取拒绝文本
            with torch.no_grad(): # 禁用梯度计算,节省内存和计算
                ref_logp_c = sequence_logprob(ref_model, tokenizer, item_prompt, item_chosen, device) # 计算参考模型选中文本的对数概率
                ref_logp_r = sequence_logprob(ref_model, tokenizer, item_prompt, item_rejected, device) # 计算参考模型拒绝文本的对数概率
            optim.zero_grad() # 清空优化器的梯度
            with torch.autocast(device_type=device, dtype=dtype): # 使用自动混合精度计算,节省内存和计算
                pi_logp_c = sequence_logprob(pi_model, tokenizer, item_prompt, item_chosen, device) # 计算策略模型选中文本的对数概率
                pi_logp_r = sequence_logprob(pi_model, tokenizer, item_prompt, item_rejected, device) # 计算策略模型拒绝文本的对数概率
                loss = dpo_loss(pi_logp_c, pi_logp_r, ref_logp_c, ref_logp_r, beta=0.1) # 使用较小的 beta
            print(f"--- Debug Info (Step {step+1}) ---")
            print(f"ref_logp_c: {ref_logp_c.item():.4f}, ref_logp_r: {ref_logp_r.item():.4f}")
            print(f"pi_logp_c: {pi_logp_c.item():.4f}, pi_logp_r: {pi_logp_r.item():.4f}")
            pi_diff_logp = pi_logp_c - pi_logp_r # 计算策略模型选中文本和拒绝文本的对数概率差
            ref_diff_logp = ref_logp_c - ref_logp_r # 计算参考模型选中文本和拒绝文本的对数概率差
            print(f"Pi Diff: {pi_diff_logp.item():.4f}, Ref Diff: {ref_diff_logp.item():.4f}")
            diff = pi_diff_logp - ref_diff_logp # 计算策略模型和参考模型的对数概率差
            print(f"Overall Diff: {diff.item():.4f}, Loss: {loss.item():.4f}")
            loss.backward() # 反向传播计算梯度
            torch.nn.utils.clip_grad_norm_(pi_model.parameters(), 1.0) # 梯度裁剪,防止梯度爆炸
            optim.step() # 更新策略模型的参数
            total_loss += loss.item() # 累计损失
            if (step + 1) % 1 == 0 or step == len(loader) - 1: # 每 1 步或最后一步打印当前平均损失
                current_avg_loss = total_loss / (step + 1) # 计算当前平均损失
                print(f"  Step {step+1}/{len(loader)}: 当前平均 DPO Loss={current_avg_loss:.4f}")
        print(f"Epoch {epoch+1} 结束: 平均 DPO Loss={total_loss/len(loader):.4f}")
    pi_model.eval() # 切换策略模型为评估模式,不进行训练
    print(f"\n{'='*20} 开始推理演示 (策略模型 vs 参考模型) {'='*20}")
    # 定义一些推理提示
    prompts_for_inference = [
        "将'你好世界'翻译成英文", "把下面句子改礼貌:把会议纪要发我",
        "一句话解释损失掩码", "用一句话描述 DPO 的作用",
    ]
    for ptxt in prompts_for_inference: # 遍历每个推理提示
        formatted_inference_prompt = QWEN_CHAT_TEMPLATE.format(instruction=ptxt) # 格式化推理提示,添加指令模板
        inputs = tokenizer(formatted_inference_prompt, return_tensors="pt").to(device) # 将格式化后的提示转换为模型输入,移动到设备上
        print(f"\n--- Prompt ---\n{ptxt}")
        # 参考模型生成
        with torch.no_grad(): # 禁用梯度计算,节省内存和计算
            gen_ref = ref_model.generate(**inputs, max_new_tokens=64, pad_token_id=tokenizer.eos_token_id) # 参考模型生成回复,设置最大新令牌数为 64,填充令牌 ID 为结束令牌 ID
        decoded_ref = tokenizer.decode(gen_ref[0][inputs.input_ids.shape[1]:], skip_special_tokens=True).strip() # 解码参考模型生成的回复,移除特殊令牌
        print(f"--- 参考模型 (Ref) 回复 ---\n{decoded_ref}")
        # 策略模型生成
        with torch.no_grad(): # 禁用梯度计算,节省内存和计算
            gen_pi = pi_model.generate(**inputs, max_new_tokens=64, pad_token_id=tokenizer.eos_token_id) # 策略模型生成回复,设置最大新令牌数为 64,填充令牌 ID 为结束令牌 ID
        decoded_pi = tokenizer.decode(gen_pi[0][inputs.input_ids.shape[1]:], skip_special_tokens=True).strip() # 解码策略模型生成的回复,移除特殊令牌
        print(f"--- 策略模型 (Pi) 回复 ---\n{decoded_pi}")
if __name__ == "__main__":
    parser = argparse.ArgumentParser(description="DPO Minimal Demo") # 解析命令行参数,设置默认值
    parser.add_argument("--model", default=os.environ.get("DPO_DEMO_MODEL", "Qwen/Qwen1.5-1.8B-Chat")) # 模型名称,默认值为 Qwen/Qwen1.5-1.8B-Chat
    parser.add_argument("--epochs", type=int, default=3) # 训练轮数,默认值为 3
    parser.add_argument("--lr", type=float, default=1e-6) # 学习率,默认值为 1e-6
    parser.add_argument("--batch_size", type=int, default=1) # 批次大小,默认值为 1
    parser.add_argument("--max_len", type=int, default=512) # 最大序列长度,默认值为 512
    parser.add_argument("--seed", type=int, default=42) # 随机种子,默认值为 42
    args = parser.parse_args() # 解析命令行参数
    main(args) # 调用主函数,传入解析后的参数

⚠️ 注意:当前聊天模板硬编码为 Qwen 格式,若更换其他模型(如 LLaMA、ChatGLM),需同步调整 QWEN_CHAT_TEMPLATE 以匹配其对话格式。

[AFFILIATE_SLOT_1]

总结与延伸思考

通过这个极简 Demo,我们可以直观感受到 DPO 的工程简洁性:仅需一个参考模型、一个策略模型与一个基于对数概率差的损失函数,即可完成人类偏好对齐任务,彻底绕开了传统 RLHF 中的奖励建模与强化学习采样环节。这种设计不仅大幅降低了训练门槛,也提升了训练稳定性。当然,生产级应用仍需考虑更大规模的偏好数据集、更精细的模板适配、分布式训练支持以及完善的评估体系。希望本文能帮助你快速掌握 DPO 的核心实现逻辑,并激发你在自己的项目中尝试这一高效对齐方法的兴趣。

[AFFILIATE_SLOT_2]