下面从零实现一个在大语言模型上使用 PPO 算法的完整系统。整体分为 四大模型 + 六大核心模块,每个模块都有完整的 PyTorch 代码、输入输出说明和数据流图。
整体架构全景
┌─────────────────────────────────────────────────────────────────┐
│ PPO 训练主循环 │
│ │
│ ┌──────────┐ ┌──────────┐ ┌──────────┐ ┌──────────┐ │
│ │ Actor │ │Reference │ │ Reward │ │ Critic │ │
│ │ (策略) │ │ (参考) │ │ (奖励) │ │ (价值) │ │
│ │ 可训练 │ │ 冻结 │ │ 冻结 │ │ 可训练 │ │
│ └────┬─────┘ └────┬─────┘ └────┬─────┘ └────┬─────┘ │
│ │ │ │ │ │
│ ▼ ▼ ▼ ▼ │
│ ┌─────────────────────────────────────────────────────────┐ │
│ │ ① Rollout → ② 奖励计算 → ③ GAE优势估计 → ④ PPO更新 │ │
│ └─────────────────────────────────────────────────────────┘ │
└─────────────────────────────────────────────────────────────────┘
模块一:数据准备
数据格式
PPO 训练需要两类数据:
python
# 类型1:Prompt 数据(用于 Rollout 生成)
# 只需用户指令,不需要回答
prompts = [
"请用Python写一个快速排序算法",
"解释一下什么是量子计算",
"帮我写一封求职邮件",
...
]
# 类型2:偏好数据(仅用于训练 Reward Model)
preference_data = [
{
"prompt": "请用Python写一个快速排序算法",
"chosen": "def quicksort(arr): ...", # 人类偏好的好回答
"rejected": "def sort(arr): arr.sort()" # 人类不喜欢的差回答
},
...
]
数据加载器实现
python
import torch
from torch.utils.data import Dataset, DataLoader
from transformers import AutoTokenizer
class PromptDataset(Dataset):
"""PPO训练用的Prompt数据集"""
def __init__(self, prompts, tokenizer, max_prompt_len=128):
self.prompts = prompts
self.tokenizer = tokenizer
self.max_prompt_len = max_prompt_len
def __len__(self):
return len(self.prompts)
def __getitem__(self, idx):
# 编码 prompt
encoded = self.tokenizer(
self.prompts[idx],
max_length=self.max_prompt_len,
truncation=True,
padding=False,
return_tensors="pt"
)
return {
"prompt_ids": encoded["input_ids"].squeeze(0), # (prompt_len,)
"prompt_mask": encoded["attention_mask"].squeeze(0), # (prompt_len,)
}
class PreferenceDataset(Dataset):
"""Reward Model训练用的偏好数据"""
def __init__(self, preference_pairs, tokenizer, max_len=512):
self.data = preference_pairs
self.tokenizer = tokenizer
self.max_len = max_len
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
item = self.data[idx]
# 拼接 prompt + chosen
chosen_text = item["prompt"] + item["chosen"]
rejected_text = item["prompt"] + item["rejected"]
chosen_tokens = self.tokenizer(
chosen_text, max_length=self.max_len,
truncation=True, padding="max_length", return_tensors="pt"
)
rejected_tokens = self.tokenizer(
rejected_text, max_length=self.max_len,
truncation=True, padding="max_length", return_tensors="pt"
)
return {
"chosen_ids": chosen_tokens["input_ids"].squeeze(0),
"chosen_mask": chosen_tokens["attention_mask"].squeeze(0),
"rejected_ids": rejected_tokens["input_ids"].squeeze(0),
"rejected_mask": rejected_tokens["attention_mask"].squeeze(0),
}
模块二:四大模型定义
Actor 模型(策略模型)--- 唯一被更新的模型
python
import torch.nn as nn
from transformers import AutoModelForCausalLM
class ActorModel(nn.Module):
"""
角色:策略模型(玩家)
作用:根据 prompt 生成 response,是训练中唯一被更新的模型
输入:prompt 的 input_ids 和 attention_mask
输出:每个 token 位置的 logits (b, l, vocab_size)
"""
def __init__(self, model_name):
super().__init__()
self.model = AutoModelForCausalLM.from_pretrained(model_name)
# 确保参数可训练
for param in self.model.parameters():
param.requires_grad = True
def forward(self, input_ids, attention_mask):
outputs = self.model(input_ids=input_ids, attention_mask=attention_mask)
return outputs.logits # (batch, seq_len, vocab_size)
def generate(self, input_ids, attention_mask, max_new_tokens=128, **kwargs):
"""自回归生成 response"""
return self.model.generate(
input_ids=input_ids,
attention_mask=attention_mask,
max_new_tokens=max_new_tokens,
do_sample=True,
temperature=1.0,
top_p=0.9,
**kwargs
)
Reference 模型(参考模型)--- 永远冻结
python
class ReferenceModel(nn.Module):
"""
角色:参考模型(初心标尺)
作用:提供 SFT 阶段的原始概率分布,用于计算 KL 散度惩罚
输入:prompt + response 的 input_ids 和 attention_mask
输出:每个 token 位置的 log 概率 (b, l)
"""
def __init__(self, model_name):
super().__init__()
self.model = AutoModelForCausalLM.from_pretrained(model_name)
# 永远冻结,不参与训练
for param in self.model.parameters():
param.requires_grad = False
self.model.eval()
@torch.no_grad()
def forward(self, input_ids, attention_mask):
outputs = self.model(input_ids=input_ids, attention_mask=attention_mask)
return outputs.logits # (batch, seq_len, vocab_size)
Reward 模型(奖励模型)--- 永远冻结
python
class RewardModel(nn.Module):
"""
角色:裁判(人类偏好的数字化身)
作用:对生成的 response 打分,分数越高代表越符合人类偏好
输入:prompt + response 的 input_ids 和 attention_mask
输出:标量奖励值 (batch,)
"""
def __init__(self, model_name):
super().__init__()
self.model = AutoModelForCausalLM.from_pretrained(model_name)
hidden_size = self.model.config.hidden_size
# 将最后的 LM head 替换为价值头(回归头)
self.reward_head = nn.Linear(hidden_size, 1, bias=False)
# 初始化:小方差,让初始评分接近 0
nn.init.normal_(self.reward_head.weight, mean=0.0, std=0.01)
# 冻结,PPO 阶段不训练
for param in self.parameters():
param.requires_grad = False
self.eval()
@torch.no_grad()
def forward(self, input_ids, attention_mask):
outputs = self.model(input_ids=input_ids, attention_mask=attention_mask)
last_hidden = outputs.last_hidden_state # (batch, seq_len, hidden)
# 取最后一个有效 token 的隐藏状态
seq_lengths = attention_mask.sum(dim=1) - 1 # (batch,)
batch_indices = torch.arange(last_hidden.size(0), device=last_hidden.device)
last_token_hidden = last_hidden[batch_indices, seq_lengths] # (batch, hidden)
# 通过回归头输出标量分数
reward = self.reward_head(last_token_hidden) # (batch, 1)
return reward.squeeze(-1) # (batch,)
Critic 模型(价值模型)--- 可训练
python
class CriticModel(nn.Module):
"""
角色:教练(价值评估师)
作用:为每个 token 位置预测"从当前位置到结束能获得的累积奖励"
输入:prompt + response 的 input_ids 和 attention_mask
输出:每个 token 位置的价值估计 (batch, seq_len)
"""
def __init__(self, model_name):
super().__init__()
self.model = AutoModelForCausalLM.from_pretrained(model_name)
hidden_size = self.model.config.hidden_size
# 价值头:hidden → 1(标量价值)
self.value_head = nn.Linear(hidden_size, 1, bias=False)
nn.init.normal_(self.value_head.weight, mean=0.0, std=0.01)
# Critic 是可训练的
for param in self.parameters():
param.requires_grad = True
def forward(self, input_ids, attention_mask):
outputs = self.model(input_ids=input_ids, attention_mask=attention_mask)
last_hidden = outputs.last_hidden_state # (batch, seq_len, hidden)
values = self.value_head(last_hidden).squeeze(-1) # (batch, seq_len)
return values
模块三:Rollout(经验采集)
这是 PPO 的第一步:用 Actor 模型生成 response,同时记录关键信息。
python
def rollout(actor_model, reference_model, tokenizer, batch_prompts,
max_prompt_len=128, max_new_tokens=128):
"""
作用:用 Actor 生成 response,并记录生成时的 log 概率和参考模型的 log 概率
输入:
- batch_prompts: 一批 prompt 文本列表
输出:
- input_ids: prompt + response 拼接后的完整序列 (batch, total_len)
- attention_mask: 注意力掩码 (batch, total_len)
- response_mask: 仅 response 部分为1,prompt部分为0 (batch, total_len)
- old_log_probs: Actor 生成时的 log 概率 (batch, response_len)
- ref_log_probs: Reference 模型的 log 概率 (batch, response_len)
"""
# 1. 编码 prompt
encoded = tokenizer(
batch_prompts,
max_length=max_prompt_len,
truncation=True,
padding="max_length",
return_tensors="pt"
)
prompt_ids = encoded["input_ids"].to(actor_model.device)
prompt_mask = encoded["attention_mask"].to(actor_model.device)
# 2. Actor 生成 response
with torch.no_grad():
generated = actor_model.generate(
input_ids=prompt_ids,
attention_mask=prompt_mask,
max_new_tokens=max_new_tokens,
do_sample=True,
temperature=1.0,
)
# 3. 拼接 prompt + response 为完整序列
input_ids = generated # (batch, prompt_len + response_len)
attention_mask = torch.ones_like(input_ids)
prompt_len = prompt_ids.size(1)
response_len = input_ids.size(1) - prompt_len
# 4. 构造 response_mask:prompt 部分为 0,response 部分为 1
response_mask = torch.zeros_like(input_ids)
response_mask[:, prompt_len:] = 1
# 5. 计算 Actor 在生成序列上的 token 级 log 概率
old_log_probs = compute_token_log_probs(actor_model, input_ids, attention_mask)
old_log_probs = old_log_probs[:, prompt_len - 1:] # 只取 response 部分
# 注意:log_prob[t] 是由 token[t-1] 预测 token[t] 的概率
# 6. 计算 Reference 模型的 token 级 log 概率
ref_log_probs = compute_token_log_probs(reference_model, input_ids, attention_mask)
ref_log_probs = ref_log_probs[:, prompt_len - 1:]
return {
"input_ids": input_ids, # (batch, total_len)
"attention_mask": attention_mask, # (batch, total_len)
"response_mask": response_mask, # (batch, total_len)
"old_log_probs": old_log_probs, # (batch, response_len)
"ref_log_probs": ref_log_probs, # (batch, response_len)
"prompt_len": prompt_len,
}
def compute_token_log_probs(model, input_ids, attention_mask):
"""
作用:计算模型在给定序列上每个 token 位置的 log 概率
输入:input_ids (batch, seq_len)
输出:log_probs (batch, seq_len - 1)
原理:用前 n-1 个 token 预测第 n 个 token 的概率
"""
logits = model(input_ids=input_ids, attention_mask=attention_mask)
if hasattr(logits, 'logits'):
logits = logits.logits
# logits: (batch, seq_len, vocab_size)
# 前 n-1 个位置的 logits 预测后 n-1 个位置的 token
shift_logits = logits[:, :-1, :] # (batch, seq_len-1, vocab_size)
shift_labels = input_ids[:, 1:] # (batch, seq_len-1)
# 转为概率再取 log
log_probs = torch.log_softmax(shift_logits, dim=-1) # (batch, seq_len-1, vocab)
# 取出实际 token 对应的 log 概率
token_log_probs = log_probs.gather(
dim=-1, index=shift_labels.unsqueeze(-1)
).squeeze(-1) # (batch, seq_len-1)
return token_log_probs
模块四:奖励计算(含 KL 散度惩罚)
python
def compute_rewards(reward_model, rollout_data, kl_coef=0.04):
"""
作用:计算每个 token 位置的奖励 = 模型打分奖励 - KL散度惩罚
输入:
- rollout_data: rollout 阶段的输出字典
- kl_coef: KL 惩罚系数 β
输出:
- rewards: 每个 token 位置的即时奖励 (batch, response_len)
原理:
R_t = r_score(仅最后一个 token 有值)- β * (log π_θ - log π_ref)
"""
input_ids = rollout_data["input_ids"]
attention_mask = rollout_data["attention_mask"]
response_mask = rollout_data["response_mask"]
old_log_probs = rollout_data["old_log_probs"]
ref_log_probs = rollout_data["ref_log_probs"]
# 1. Reward Model 对整个序列打分(标量)
with torch.no_grad():
sequence_reward = reward_model(input_ids, attention_mask) # (batch,)
batch_size, response_len = old_log_probs.shape
# 2. 计算 token 级 KL 散度惩罚
# KL ≈ log π_θ(a|s) - log π_ref(a|s)
kl_penalty = old_log_probs - ref_log_probs # (batch, response_len)
# 3. 构造 token 级奖励
# 默认:每个 token 的奖励 = -β * KL(防止偏离参考模型)
rewards = -kl_coef * kl_penalty # (batch, response_len)
# 4. 在最后一个有效 token 位置加上序列级奖励
# 找到每个序列的最后一个有效 response token
response_lengths = response_mask.sum(dim=1) # (batch,)
for i in range(batch_size):
last_idx = int(response_lengths[i].item()) - 1
if last_idx >= 0:
rewards[i, last_idx] += sequence_reward[i]
return rewards # (batch, response_len)
模块五:GAE 优势估计
python
def compute_gae(rewards, values, response_mask, gamma=0.99, lam=0.95):
"""
作用:计算广义优势估计(Generalized Advantage Estimation)
输入:
- rewards: 每个 token 的即时奖励 (batch, response_len)
- values: Critic 预测的每个 token 的价值 (batch, response_len)
- response_mask: 有效 token 掩码 (batch, response_len)
- gamma: 折扣因子,衡量未来奖励的重要性
- lam: GAE 的 λ 参数,控制偏差-方差权衡
输出:
- advantages: 优势值 (batch, response_len)
- returns: 回报值 = advantages + values (batch, response_len)
原理:
δ_t = R_t + γ * V(s_{t+1}) - V(s_t) ← TD 误差
A_t = Σ_{l=0}^{T-t} (γλ)^l * δ_{t+l} ← 加权累积 TD 误差
"""
batch_size, seq_len = rewards.shape
advantages = torch.zeros_like(rewards)
# 从后往前递推计算
last_gae = torch.zeros(batch_size, device=rewards.device)
for t in reversed(range(seq_len)):
# 获取 t+1 时刻的价值(最后一个 token 的 V(s_{t+1}) = 0)
if t == seq_len - 1:
next_values = torch.zeros(batch_size, device=values.device)
else:
next_values = values[:, t + 1]
# 当前时刻的价值
current_values = values[:, t]
# TD 误差: δ_t = R_t + γ * V(s_{t+1}) - V(s_t)
delta = rewards[:, t] + gamma * next_values - current_values
# GAE 递推: A_t = δ_t + γλ * A_{t+1}
last_gae = delta + gamma * lam * last_gae
advantages[:, t] = last_gae
# 回报 = 优势 + 价值
returns = advantages + values
# 应用掩码:只对 response 部分计算优势
advantages = advantages * response_mask
returns = returns * response_mask
# 优势标准化(稳定训练)
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
return advantages, returns
模块六:PPO 损失计算与参数更新
python
def compute_ppo_loss(actor_model, old_log_probs, advantages, returns,
values, rollout_data, clip_epsilon=0.2,
vf_coef=0.5, ent_coef=0.01):
"""
作用:计算 PPO 的总损失 = 策略损失 + 价值损失 - 熵正则
输入:
- actor_model: 当前策略模型
- old_log_probs: 生成时的 log 概率 (batch, response_len)
- advantages: GAE 优势值 (batch, response_len)
- returns: 回报值 (batch, response_len)
- values: Critic 预测的价值 (batch, response_len)
- rollout_data: rollout 数据
- clip_epsilon: PPO clip 范围 ε
- vf_coef: 价值损失的权重
- ent_coef: 熵正则的权重(鼓励探索)
输出:
- total_loss: 总损失标量
"""
input_ids = rollout_data["input_ids"]
attention_mask = rollout_data["attention_mask"]
response_mask = rollout_data["response_mask"]
prompt_len = rollout_data["prompt_len"]
# 1. 用当前 Actor 重新计算 log 概率
current_log_probs = compute_token_log_probs(actor_model, input_ids, attention_mask)
current_log_probs = current_log_probs[:, prompt_len - 1:] # 只取 response 部分
# 2. 计算重要性采样比率 r(θ) = π_θ / π_θ_old
ratio = torch.exp(current_log_probs - old_log_probs) # (batch, response_len)
# 3. PPO-Clip 策略损失
# L_clip = min(ratio * A, clip(ratio, 1-ε, 1+ε) * A)
surr1 = ratio * advantages
surr2 = torch.clamp(ratio, 1.0 - clip_epsilon, 1.0 + clip_epsilon) * advantages
policy_loss = -torch.min(surr1, surr2) # 取负号因为是最大化目标
# 只对 response 部分求平均
policy_loss = (policy_loss * response_mask).sum() / response_mask.sum()
# 4. 价值函数损失(MSE)
value_loss = ((returns - values) ** 2) * response_mask
value_loss = value_loss.sum() / response_mask.sum()
# 5. 熵正则(鼓励探索,防止策略过早收敛)
logits = actor_model(input_ids=input_ids, attention_mask=attention_mask).logits
logits = logits[:, prompt_len - 1:-1, :] # 对齐 response 部分
probs = torch.softmax(logits, dim=-1)
log_probs_dist = torch.log_softmax(logits, dim=-1)
entropy = -(probs * log_probs_dist).sum(dim=-1) # (batch, response_len)
entropy = (entropy * response_mask).sum() / response_mask.sum()
# 6. 总损失
total_loss = policy_loss + vf_coef * value_loss - ent_coef * entropy
return total_loss, policy_loss.item(), value_loss.item(), entropy.item()
def ppo_update(actor_model, critic_model, actor_optimizer, critic_optimizer,
rollout_data, rewards, num_ppo_epochs=4):
"""
作用:执行多轮 PPO 更新(Actor + Critic)
输入:
- rollout_data: rollout 阶段采集的数据
- rewards: token 级奖励 (batch, response_len)
- num_ppo_epochs: PPO 内部更新轮数
"""
prompt_len = rollout_data["prompt_len"]
input_ids = rollout_data["input_ids"]
attention_mask = rollout_data["attention_mask"]
response_mask = rollout_data["response_mask"]
old_log_probs = rollout_data["old_log_probs"]
for epoch in range(num_ppo_epochs):
# 1. Critic 预测价值
values = critic_model(input_ids, attention_mask)
values = values[:, prompt_len - 1:-1] # 对齐 response 部分
# 2. 计算 GAE 优势
advantages, returns = compute_gae(rewards, values, response_mask)
# 3. 计算 PPO 损失
total_loss, policy_loss, value_loss, entropy = compute_ppo_loss(
actor_model, old_log_probs, advantages, returns, values, rollout_data
)
# 4. 更新 Actor
actor_optimizer.zero_grad()
total_loss.backward()
torch.nn.utils.clip_grad_norm_(actor_model.parameters(), max_norm=1.0)
actor_optimizer.step()
# 5. 单独更新 Critic(用 returns 作为目标)
with torch.no_grad():
new_values = critic_model(input_ids, attention_mask)
new_values = new_values[:, prompt_len - 1:-1]
critic_loss = ((returns - new_values) ** 2 * response_mask).sum() / response_mask.sum()
critic_optimizer.zero_grad()
critic_loss.backward()
torch.nn.utils.clip_grad_norm_(critic_model.parameters(), max_norm=1.0)
critic_optimizer.step()
return {
"total_loss": total_loss.item(),
"policy_loss": policy_loss,
"value_loss": value_loss,
"entropy": entropy,
}
完整训练主循环
python
def train_ppo(
actor_model,
reference_model,
reward_model,
critic_model,
tokenizer,
prompts,
batch_size=4,
max_prompt_len=128,
max_new_tokens=64,
kl_coef=0.04,
num_episodes=100,
lr_actor=1e-6,
lr_critic=5e-6,
):
"""
PPO 完整训练主循环
"""
actor_optimizer = torch.optim.AdamW(actor_model.parameters(), lr=lr_actor)
critic_optimizer = torch.optim.AdamW(critic_model.parameters(), lr=lr_critic)
for episode in range(num_episodes):
# ---- Step 1: 采样 prompt 批次 ----
batch_indices = torch.randint(0, len(prompts), (batch_size,))
batch_prompts = [prompts[i] for i in batch_indices]
# ---- Step 2: Rollout(经验采集)----
rollout_data = rollout(
actor_model, reference_model, tokenizer,
batch_prompts, max_prompt_len, max_new_tokens
)
# ---- Step 3: 计算奖励(含 KL 惩罚)----
rewards = compute_rewards(reward_model, rollout_data, kl_coef=kl_coef)
# ---- Step 4: PPO 更新(Actor + Critic)----
metrics = ppo_update(
actor_model, critic_model,
actor_optimizer, critic_optimizer,
rollout_data, rewards,
num_ppo_epochs=4
)
# ---- 日志 ----
if episode % 10 == 0:
avg_reward = rewards.sum(dim=1).mean().item()
print(f"Episode {episode} | "
f"Avg Reward: {avg_reward:.4f} | "
f"Policy Loss: {metrics['policy_loss']:.4f} | "
f"Value Loss: {metrics['value_loss']:.4f} | "
f"Entropy: {metrics['entropy']:.4f}")
# ============================================================
# 启动训练
# ============================================================
if __name__ == "__main__":
from transformers import AutoTokenizer
model_name = "gpt2" # 用小模型演示,实际替换为 LLaMA 等
tokenizer = AutoTokenizer.from_pretrained(model_name)
tokenizer.pad_token = tokenizer.eos_token
device = "cuda" if torch.cuda.is_available() else "cpu"
# 初始化四大模型
actor = ActorModel(model_name).to(device)
reference = ReferenceModel(model_name).to(device)
reward = RewardModel(model_name).to(device)
critic = CriticModel(model_name).to(device)
# 模拟 prompt 数据
prompts = [
"Explain quantum computing in simple terms.",
"Write a short poem about the ocean.",
"What are the benefits of exercise?",
"How does photosynthesis work?",
"Describe the history of the internet.",
] * 20 # 重复以增加数据量
# 开始 PPO 训练
train_ppo(
actor_model=actor,
reference_model=reference,
reward_model=reward,
critic_model=critic,
tokenizer=tokenizer,
prompts=prompts,
batch_size=4,
num_episodes=100,
)
数据流维度变化总结
以 batch_size=4, prompt_len=32, response_len=64 为例:
Prompt 文本 (4 条)
│
▼ Tokenizer
input_ids (4, 32)
│
▼ Actor.generate()
generated_ids (4, 96) ← prompt(32) + response(64)
│
▼ compute_token_log_probs
old_log_probs (4, 64) ← response 部分每个 token 的 log 概率
ref_log_probs (4, 64)
│
▼ Reward Model 打分
sequence_reward (4,) ← 整个序列一个标量
│
▼ 构造 token 级奖励
rewards (4, 64) ← 仅最后一个 token 有 r_score,其余为 -β*KL
│
▼ Critic 预测价值
values (4, 64) ← 每个 token 位置的 V(s)
│
▼ GAE 计算
advantages (4, 64) ← 优势值
returns (4, 64) ← 回报值
│
▼ PPO Loss
policy_loss + value_loss - entropy → 标量 → backward → 更新 Actor + Critic
各模块职责速查表
| 模块 | 模型 | 是否训练 | 核心作用 |
|---|---|---|---|
| Rollout | Actor | ✅ | 生成 response,记录 log 概率 |
| 奖励计算 | Reward + Reference | ❌ | 打分 + KL 惩罚,构造 token 级奖励 |
| GAE | Critic | ✅ | 估计每个 token 的优势值 |
| PPO 更新 | Actor + Critic | ✅ | clip 策略梯度 + 价值回归 |
| Reference | --- | ❌ | 提供基准分布,防止策略跑偏 |
整个系统的核心思想可以概括为一句话:Actor 负责"说",Reward 负责"打分",Reference 负责"纠偏",Critic 负责"指导",PPO 负责"安全地学"。