在大语言模型上使用 PPO 算法

下面从零实现一个在大语言模型上使用 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 负责"安全地学"。

相关推荐
studyrunner1 小时前
【AI开源】Buzz 实战教程:搭建人类与多 AI Agent 协同工作的自托管工作区
人工智能·开源
动物园猫1 小时前
夜间野生动物目标检测数据集:17类别、17,000张图像 | 目标检测
人工智能·目标检测·计算机视觉
小马9261 小时前
2026年8月4日科技热点深度解析:AI大模型群雄逐鹿、卫星互联网组网提速、半导体封装材料革命
人工智能·科技·deepseek
IT_陈寒2 小时前
React的useEffect依赖数组把我坑惨了,原来这样写才靠谱
前端·人工智能·后端
盖伦发发2 小时前
AIE-AI Engineering三, 四章总结: 如何评估AI应用
人工智能·ai
大模型搬砖师2 小时前
在Kubernetes上部署企业AI网关:一份云原生参考
网络·人工智能·安全
love530love2 小时前
Ubuntu系统通过Homebrew安装Lightpanda完整实战教程(含端口占用排坑)
大数据·linux·运维·人工智能·elasticsearch·搜索引擎
V哥AI增长2 小时前
ChatGPT/Perplexity引用机制解析:AI搜索引擎的语义解析与信任评估体系
人工智能·搜索引擎·chatgpt
zander2582 小时前
35. 搜索插入位置:从边界语义理解二分查找
数据结构·算法·leetcode