RLHF人类反馈强化学习深度解析:从偏好数据到对齐策略的完整工程实现

举报
柠檬🍋 发表于 2026/08/16 03:39:02 2026/08/16
【摘要】 RLHF人类反馈强化学习深度解析:从偏好数据到对齐策略的完整工程实现 一、引言:预训练模型为何需要对齐大语言模型经过预训练后具备了语言生成能力,但预训练目标是最大化下一个token的概率,并不保证模型输出符合人类价值观和用户意图。预训练模型可能出现有害输出、幻觉、不遵循指令等问题。RLHF(Reinforcement Learning from Human Feedback)通过收集人类偏...

RLHF人类反馈强化学习深度解析:从偏好数据到对齐策略的完整工程实现

一、引言:预训练模型为何需要对齐

大语言模型经过预训练后具备了语言生成能力,但预训练目标是最大化下一个token的概率,并不保证模型输出符合人类价值观和用户意图。预训练模型可能出现有害输出、幻觉、不遵循指令等问题。RLHF(Reinforcement Learning from Human Feedback)通过收集人类偏好数据训练奖励模型,再用强化学习优化模型策略,使模型输出与人类偏好对齐。这一方法被ChatGPT采用并取得突破性效果,成为现代大模型对齐的标准流程。本文将完整解析RLHF的三阶段流程:监督微调(SFT)、奖励模型训练和PPO强化学习优化,并提供可运行的PyTorch实现。

二、RLHF三阶段流程概述

RLHF包含三个核心阶段。第一阶段是监督微调(Supervised Fine-Tuning, SFT),使用人工编写的高质量指令-回复对微调预训练模型,使模型学会遵循指令的基本格式。第二阶段是奖励模型(Reward Model, RM)训练,收集人类对模型多个输出的偏好排序数据,训练一个能够预测人类偏好的标量奖励模型。第三阶段是强化学习优化,使用PPO(Proximal Policy Optimization)算法,以奖励模型的输出作为奖励信号,优化SFT模型使其生成更高奖励的回复,同时通过KL散度惩罚防止模型偏离原始分布过远。

import torch
import torch.nn as nn
import torch.nn.functional as F
import math
from typing import List, Dict, Tuple, Optional
from dataclasses import dataclass, field
import random
import os

@dataclass
class SFTExample:
    """监督微调数据样例"""
    instruction: str
    input: str = ""
    output: str = ""

@dataclass
class PreferenceExample:
    """偏好数据样例(用于奖励模型训练)"""
    prompt: str
    response_chosen: str  # 人类偏好的回复
    response_rejected: str  # 被拒绝的回复
    score_chosen: float = 1.0
    score_rejected: float = 0.0

@dataclass
class RLHFConfig:
    """RLHF配置"""
    # 模型配置
    vocab_size: int = 32000
    dim: int = 256
    n_layers: int = 4
    n_heads: int = 8
    
    # SFT配置
    sft_learning_rate: float = 2e-5
    sft_epochs: int = 3
    
    # 奖励模型配置
    rm_learning_rate: float = 5e-5
    rm_epochs: int = 2
    
    # PPO配置
    ppo_learning_rate: float = 1e-5
    ppo_epochs: int = 4
    ppo_clip: float = 0.2
    kl_coeff: float = 0.1
    reward_clip: float = 5.0
    max_response_length: int = 256
    
    # 训练配置
    batch_size: int = 4
    device: str = 'cuda' if torch.cuda.is_available() else 'cpu'


class SimpleLanguageModel(nn.Module):
    """简化的大语言模型(用于演示RLHF流程)"""
    
    def __init__(self, vocab_size=32000, dim=256, n_layers=4, n_heads=8):
        super().__init__()
        self.vocab_size = vocab_size
        self.dim = dim
        
        self.embedding = nn.Embedding(vocab_size, dim)
        self.pos_embedding = nn.Embedding(512, dim)
        
        # 简化的Transformer层
        self.layers = nn.ModuleList([
            nn.TransformerEncoderLayer(
                d_model=dim, nhead=n_heads,
                dim_feedforward=dim * 4,
                dropout=0.1, batch_first=True,
                norm_first=True  # Pre-LN
            )
            for _ in range(n_layers)
        ])
        self.norm = nn.LayerNorm(dim)
        self.lm_head = nn.Linear(dim, vocab_size, bias=False)
        
        # 权重共享
        self.lm_head.weight = self.embedding.weight
        
        self._init_weights()
    
    def _init_weights(self):
        for m in self.modules():
            if isinstance(m, nn.Linear):
                nn.init.normal_(m.weight, mean=0, std=0.02)
                if m.bias is not None:
                    nn.init.zeros_(m.bias)
            elif isinstance(m, nn.Embedding):
                nn.init.normal_(m.weight, mean=0, std=0.02)
    
    def forward(self, input_ids, return_hidden=False):
        B, T = input_ids.shape
        pos = torch.arange(T, device=input_ids.device).unsqueeze(0).expand(B, T)
        
        x = self.embedding(input_ids) + self.pos_embedding(pos)
        
        # 创建因果掩码
        causal_mask = torch.triu(torch.ones(T, T), diagonal=1).bool()
        for layer in self.layers:
            x = layer(x, src_key_padding_mask=None, 
                     mask=~causal_mask.to(x.device) if False else None,
                     is_causal=True)
        
        x = self.norm(x)
        
        if return_hidden:
            return self.lm_head(x), x
        return self.lm_head(x)
    
    @torch.no_grad()
    def generate(self, input_ids, max_length=100, temperature=0.7, 
                 top_p=0.9, do_sample=True):
        """生成文本"""
        self.eval()
        for _ in range(max_length - input_ids.size(1)):
            logits = self.forward(input_ids)
            next_logits = logits[:, -1, :] / temperature
            
            if do_sample:
                # Top-p采样
                sorted_logits, sorted_indices = torch.sort(next_logits, descending=True)
                cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
                sorted_indices_to_remove = cumulative_probs > top_p
                sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
                sorted_indices_to_remove[..., 0] = 0
                indices_to_remove = sorted_indices_to_remove.scatter(
                    1, sorted_indices, sorted_indices_to_remove)
                next_logits = next_logits.masked_fill(indices_to_remove, float('-inf'))
                
                probs = F.softmax(next_logits, dim=-1)
                next_token = torch.multinomial(probs, num_samples=1)
            else:
                next_token = next_logits.argmax(dim=-1, keepdim=True)
            
            input_ids = torch.cat([input_ids, next_token], dim=1)
            
            # 停止条件(简化:遇到EOS或达到最大长度)
            if next_token.item() == 2:  # EOS
                break
        
        return input_ids


def test_base_model():
    """测试基础模型"""
    model = SimpleLanguageModel(vocab_size=1000, dim=128, n_layers=2, n_heads=4)
    
    x = torch.randint(0, 1000, (2, 32))
    logits = model(x)
    print(f"输入: {x.shape}")
    print(f"输出: {logits.shape}")
    
    # 生成
    prompt = torch.randint(0, 1000, (1, 10))
    generated = model.generate(prompt, max_length=20, do_sample=True)
    print(f"生成: {generated.shape}")
    
    params = sum(p.numel() for p in model.parameters())
    print(f"模型参数: {params:,}")

if __name__ == "__main__":
    test_base_model()

三、第一阶段:监督微调(SFT)

SFT阶段使用人工标注的指令-回复对数据微调预训练模型。数据格式通常为Alpaca格式:instruction(指令)、input(可选输入)、output(期望输出)。训练时将instruction和output拼接,对instruction部分不计算损失(mask掉),只对output部分计算语言模型损失。

class SFTTrainer:
    """监督微调训练器"""
    
    def __init__(self, model: nn.Module, config: RLHFConfig):
        self.model = model
        self.config = config
        self.device = config.device
        
        self.model.to(self.device)
        
        self.optimizer = torch.optim.AdamW(
            model.parameters(),
            lr=config.sft_learning_rate,
            weight_decay=0.01,
            betas=(0.9, 0.999)
        )
    
    def format_prompt(self, example: SFTExample) -> str:
        """格式化指令"""
        if example.input:
            return (f"Below is an instruction that describes a task, "
                    f"paired with an input that provides further context. "
                    f"Write a response that appropriately completes the request.\n\n"
                    f"### Instruction:\n{example.instruction}\n\n"
                    f"### Input:\n{example.input}\n\n"
                    f"### Response:\n{example.output}")
        else:
            return (f"Below is an instruction that describes a task. "
                    f"Write a response that appropriately completes the request.\n\n"
                    f"### Instruction:\n{example.instruction}\n\n"
                    f"### Response:\n{example.output}")
    
    def prepare_batch(self, examples: List[SFTExample], tokenizer) -> Tuple[torch.Tensor, torch.Tensor]:
        """准备训练批次"""
        max_len = 256
        input_ids_list = []
        labels_list = []
        
        for ex in examples:
            # 格式化但不包含output
            prompt = self.format_prompt(ex).replace(ex.output, "")
            response = ex.output
            
            # 简化的tokenization(实际使用真实tokenizer)
            prompt_ids = tokenizer.encode(prompt)
            response_ids = tokenizer.encode(response) + [2]  # EOS
            
            input_ids = prompt_ids + response_ids
            labels = [-100] * len(prompt_ids) + response_ids  # prompt部分不计算损失
            
            # 截断/填充
            if len(input_ids) > max_len:
                input_ids = input_ids[:max_len]
                labels = labels[:max_len]
            else:
                padding = max_len - len(input_ids)
                input_ids = input_ids + [0] * padding
                labels = labels + [-100] * padding
            
            input_ids_list.append(input_ids)
            labels_list.append(labels)
        
        return (
            torch.tensor(input_ids_list, device=self.device),
            torch.tensor(labels_list, device=self.device)
        )
    
    def train_step(self, input_ids, labels):
        """训练一步"""
        self.model.train()
        
        logits = self.model(input_ids)
        
        # 计算损失(忽略-100的标签)
        loss = F.cross_entropy(
            logits.view(-1, logits.size(-1)),
            labels.view(-1),
            ignore_index=-100
        )
        
        self.optimizer.zero_grad()
        loss.backward()
        torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)
        self.optimizer.step()
        
        return loss.item()
    
    def train(self, examples: List[SFTExample], tokenizer, epochs: int = 3):
        """完整训练"""
        batch_size = self.config.batch_size
        
        for epoch in range(epochs):
            random.shuffle(examples)
            total_loss = 0
            n_batches = 0
            
            for i in range(0, len(examples), batch_size):
                batch = examples[i:i+batch_size]
                input_ids, labels = self.prepare_batch(batch, tokenizer)
                loss = self.train_step(input_ids, labels)
                total_loss += loss
                n_batches += 1
                
                if n_batches % 10 == 0:
                    print(f"Epoch {epoch+1}, Batch {n_batches}, Loss: {loss:.4f}")
            
            avg_loss = total_loss / n_batches
            print(f"Epoch {epoch+1} 完成, 平均Loss: {avg_loss:.4f}")


# 简单的tokenizer模拟
class SimpleTokenizer:
    def __init__(self, vocab_size=32000):
        self.vocab_size = vocab_size
        self.char_to_id = {}
        self.id_to_char = {}
        # 特殊token
        for i, tok in enumerate(['<pad>', '<unk>', '<eos>', '<bos>']):
            self.char_to_id[tok] = i
            self.id_to_char[i] = tok
    
    def encode(self, text: str) -> List[int]:
        ids = []
        for char in text:
            if char not in self.char_to_id:
                idx = len(self.char_to_id) + hash(char) % (self.vocab_size - 100)
                self.char_to_id[char] = idx
                self.id_to_char[idx] = char
            ids.append(self.char_to_id[char])
        return ids
    
    def decode(self, ids: List[int]) -> str:
        return ''.join([self.id_to_char.get(i, '<unk>') for i in ids])


def test_sft():
    """测试SFT训练"""
    config = RLHFConfig(dim=128, n_layers=2, n_heads=4, batch_size=4)
    model = SimpleLanguageModel(vocab_size=1000, dim=128, n_layers=2, n_heads=4)
    
    trainer = SFTTrainer(model, config)
    tokenizer = SimpleTokenizer(vocab_size=1000)
    
    # 模拟SFT数据
    examples = [
        SFTExample(instruction="What is AI?", output="AI is artificial intelligence."),
        SFTExample(instruction="Explain deep learning.", output="Deep learning uses neural networks."),
        SFTExample(instruction="What is NLP?", output="NLP is natural language processing."),
    ] * 10
    
    trainer.train(examples, tokenizer, epochs=2)

if __name__ == "__main__":
    test_sft()

四、第二阶段:奖励模型训练

奖励模型接收prompt和response作为输入,输出一个标量奖励值。训练数据为人类偏好对(chosen和rejected),使用Bradley-Terry模型定义偏好概率:

P(ywylx)=σ(r(x,yw)r(x,yl))P(y_w \succ y_l | x) = \sigma(r(x, y_w) - r(x, y_l))

其中 r(x,y)r(x, y) 是奖励模型对输入 xx、回复 yy 的评分,ywy_w 是偏好回复,yly_l 是被拒绝回复,σ\sigma 是sigmoid函数。损失函数为负对数似然:

L=logσ(r(x,yw)r(x,yl))\mathcal{L} = -\log\sigma(r(x, y_w) - r(x, y_l))

class RewardModel(nn.Module):
    """奖励模型"""
    
    def __init__(self, base_model: SimpleLanguageModel):
        super().__init__()
        # 复用基础模型的结构(实际中加载SFT模型权重)
        self.model = base_model
        # 奖励头:将最后的hidden state映射为标量
        self.reward_head = nn.Linear(base_model.dim, 1)
        
        # 冻结基础模型(可选)
        for param in self.model.parameters():
            param.requires_grad = True  # 奖励模型通常也训练所有参数
    
        # 初始化奖励头
        nn.init.normal_(self.reward_head.weight, mean=0, std=0.02)
        nn.init.zeros_(self.reward_head.bias)
    
    def forward(self, input_ids: torch.Tensor) -> torch.Tensor:
        """计算奖励值"""
        # 获取最后一层的hidden state
        logits, hidden = self.model(input_ids, return_hidden=True)
        
        # 使用最后一个token的hidden state计算奖励
        # 实际中也可以使用均值池化或特定位置的token
        last_hidden = hidden[:, -1, :]  # (batch, dim)
        reward = self.reward_head(last_hidden)  # (batch, 1)
        
        return reward.squeeze(-1)  # (batch,)
    
    def compute_pairwise_loss(
        self, 
        chosen_ids: torch.Tensor, 
        rejected_ids: torch.Tensor
    ) -> Tuple[torch.Tensor, float]:
        """计算偏好对损失"""
        chosen_rewards = self.forward(chosen_ids)  # (batch,)
        rejected_rewards = self.forward(rejected_ids)  # (batch,)
        
        # Bradley-Terry模型: -log(sigmoid(r_chosen - r_rejected))
        logits = chosen_rewards - rejected_rewards
        loss = -F.logsigmoid(logits).mean()
        
        # 计算准确率
        with torch.no_grad():
            accuracy = (logits > 0).float().mean().item()
        
        return loss, accuracy


class RewardModelTrainer:
    """奖励模型训练器"""
    
    def __init__(self, reward_model: RewardModel, config: RLHFConfig):
        self.reward_model = reward_model
        self.config = config
        self.device = config.device
        
        self.reward_model.to(self.device)
        
        self.optimizer = torch.optim.AdamW(
            reward_model.parameters(),
            lr=config.rm_learning_rate,
            weight_decay=0.01,
            betas=(0.9, 0.999)
        )
    
    def prepare_batch(self, examples: List[PreferenceExample], tokenizer):
        """准备训练批次"""
        max_len = self.config.max_response_length
        chosen_ids_list = []
        rejected_ids_list = []
        
        for ex in examples:
            chosen_text = ex.prompt + ex.response_chosen
            rejected_text = ex.prompt + ex.response_rejected
            
            chosen_ids = tokenizer.encode(chosen_text)
            rejected_ids = tokenizer.encode(rejected_text)
            
            # 填充到相同长度
            max_len_pair = max(len(chosen_ids), len(rejected_ids))
            max_len_pair = min(max_len_pair, max_len)
            
            chosen_ids = chosen_ids[:max_len_pair] + [0] * (max_len_pair - len(chosen_ids[:max_len_pair]))
            rejected_ids = rejected_ids[:max_len_pair] + [0] * (max_len_pair - len(rejected_ids[:max_len_pair]))
            
            chosen_ids_list.append(chosen_ids)
            rejected_ids_list.append(rejected_ids)
        
        return (
            torch.tensor(chosen_ids_list, device=self.device),
            torch.tensor(rejected_ids_list, device=self.device)
        )
    
    def train_step(self, chosen_ids, rejected_ids):
        """训练一步"""
        self.reward_model.train()
        
        loss, accuracy = self.reward_model.compute_pairwise_loss(chosen_ids, rejected_ids)
        
        self.optimizer.zero_grad()
        loss.backward()
        torch.nn.utils.clip_grad_norm_(self.reward_model.parameters(), 1.0)
        self.optimizer.step()
        
        return loss.item(), accuracy
    
    def train(self, examples: List[PreferenceExample], tokenizer, epochs: int = 2):
        """完整训练"""
        batch_size = self.config.batch_size
        
        for epoch in range(epochs):
            random.shuffle(examples)
            total_loss = 0
            total_acc = 0
            n_batches = 0
            
            for i in range(0, len(examples), batch_size):
                batch = examples[i:i+batch_size]
                chosen_ids, rejected_ids = self.prepare_batch(batch, tokenizer)
                loss, acc = self.train_step(chosen_ids, rejected_ids)
                total_loss += loss
                total_acc += acc
                n_batches += 1
                
                if n_batches % 10 == 0:
                    print(f"Epoch {epoch+1}, Batch {n_batches}, "
                          f"Loss: {loss:.4f}, Acc: {acc:.2%}")
            
            print(f"Epoch {epoch+1} 完成, 平均Loss: {total_loss/n_batches:.4f}, "
                  f"平均Acc: {total_acc/n_batches:.2%}")


def test_reward_model():
    """测试奖励模型"""
    config = RLHFConfig(dim=128, n_layers=2, n_heads=4, batch_size=4)
    base_model = SimpleLanguageModel(vocab_size=1000, dim=128, n_layers=2, n_heads=4)
    reward_model = RewardModel(base_model)
    trainer = RewardModelTrainer(reward_model, config)
    tokenizer = SimpleTokenizer(vocab_size=1000)
    
    # 模拟偏好数据
    examples = [
        PreferenceExample(
            prompt="What is AI?",
            response_chosen="AI is artificial intelligence, a branch of computer science.",
            response_rejected="AI is bad."
        ),
        PreferenceExample(
            prompt="Explain ML.",
            response_chosen="Machine learning is a method of data analysis.",
            response_rejected="I don't know."
        ),
    ] * 10
    
    trainer.train(examples, tokenizer, epochs=3)
    
    # 测试奖励值
    test_input = torch.randint(0, 1000, (4, 32))
    rewards = reward_model(test_input)
    print(f"\n测试奖励值: {rewards.tolist()}")

if __name__ == "__main__":
    test_reward_model()

奖励模型训练中的一个重要问题是奖励值分布的偏移。随着训练进行,奖励值的绝对大小会不断增大(模型倾向于给出越来越高的分数),这会导致PPO阶段的不稳定。解决方案包括:使用偏好对训练而非绝对评分、对奖励值进行标准化、限制奖励值的范围。此外,奖励模型的泛化能力也是一个挑战——模型可能在训练数据分布上表现良好但对分布外的输入给出错误评分,这需要多样化的偏好数据来缓解。

五、第三阶段:PPO强化学习优化

PPO是RLHF最核心也最复杂的阶段。训练流程为:对每个prompt,使用当前策略模型生成回复,计算旧策略下的对数概率;使用奖励模型对生成的回复评分;使用优势估计(Advantage Estimation)和PPO的clip目标函数更新策略模型;同时通过KL散度惩罚防止策略偏离参考模型太远。

PPO的clip目标函数为:

LCLIP(θ)=E^t[min(rt(θ)A^t,clip(rt(θ),1ϵ,1+ϵ)A^t)]L^{CLIP}(\theta) = \hat{E}_t[\min(r_t(\theta)\hat{A}_t, \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon)\hat{A}_t)]

其中 rt(θ)=πθ(atst)πold(atst)r_t(\theta) = \frac{\pi_\theta(a_t|s_t)}{\pi_{old}(a_t|s_t)} 是新旧策略的概率比,A^t\hat{A}_t 是优势估计,ϵ\epsilon 是裁剪范围。

class PPOTrainer:
    """PPO训练器"""
    
    def __init__(
        self,
        policy_model: SimpleLanguageModel,  # 要优化的策略模型
        reference_model: SimpleLanguageModel,  # 参考模型(冻结)
        reward_model: RewardModel,  # 奖励模型(冻结)
        config: RLHFConfig
    ):
        self.policy = policy_model
        self.reference = reference_model
        self.reward_model = reward_model
        self.config = config
        self.device = config.device
        
        self.policy.to(self.device)
        self.reference.to(self.device)
        self.reward_model.to(self.device)
        
        # 冻结参考模型和奖励模型
        for param in self.reference.parameters():
            param.requires_grad = False
        for param in self.reward_model.parameters():
            param.requires_grad = False
        
        # PPO优化器
        self.optimizer = torch.optim.AdamW(
            self.policy.parameters(),
            lr=config.ppo_learning_rate,
            betas=(0.9, 0.999),
            weight_decay=0.0
        )
        
        # 价值模型(用于优势估计,实际中通常共享策略模型的部分参数)
        self.value_head = nn.Linear(config.dim, 1).to(self.device)
        self.value_optimizer = torch.optim.AdamW(
            self.value_head.parameters(),
            lr=config.ppo_learning_rate
        )
    
    def get_log_probs(self, model, input_ids, labels):
        """计算模型在给定序列上的对数概率"""
        logits = model(input_ids)
        log_probs = F.log_softmax(logits, dim=-1)
        
        # 收集对应label的对数概率
        # 对齐:logits[:, :-1] 对应 labels[:, 1:]
        log_probs_labels = log_probs[:, :-1, :].gather(
            -1, labels[:, 1:].unsqueeze(-1)
        ).squeeze(-1)
        
        return log_probs_labels
    
    def compute_kl_penalty(self, policy_log_probs, ref_log_probs):
        """计算KL散度惩罚"""
        # KL(p || q) = sum p * (log p - log q)
        kl = (policy_log_probs - ref_log_probs).mean()
        return kl
    
    def generate_responses(self, prompts, max_length=64):
        """使用策略模型生成回复"""
        self.policy.eval()
        responses = []
        
        with torch.no_grad():
            for prompt in prompts:
                # 生成
                generated = self.policy.generate(
                    prompt.unsqueeze(0),
                    max_length=max_length,
                    temperature=0.7,
                    do_sample=True
                )
                responses.append(generated)
        
        # 填充到相同长度
        max_len = max(r.size(1) for r in responses)
        padded = []
        for r in responses:
            if r.size(1) < max_len:
                padding = torch.full(
                    (1, max_len - r.size(1)), 0,
                    dtype=torch.long, device=self.device
                )
                r = torch.cat([r, padding], dim=1)
            padded.append(r)
        
        return torch.cat(padded, dim=0)
    
    def compute_rewards(self, responses, prompt_lengths):
        """计算奖励值(奖励模型 + KL惩罚)"""
        with torch.no_grad():
            # 基础奖励
            base_rewards = self.reward_model(responses)  # (batch,)
            
            # KL惩罚
            policy_log_probs = self.get_log_probs(self.policy, responses, responses)
            ref_log_probs = self.get_log_probs(self.reference, responses, responses)
            kl = policy_log_probs - ref_log_probs  # 每个token的KL
            
            # 将KL惩罚加到奖励上
            # 通常在最后一个token施加惩罚
            rewards = base_rewards.clone()
            for i in range(len(rewards)):
                actual_len = prompt_lengths[i] + 1  # 简化
                kl_penalty = -self.config.kl_coeff * kl[i, :actual_len].mean()
                rewards[i] += kl_penalty
            
            # 奖励裁剪
            rewards = rewards.clamp(-self.config.reward_clip, self.config.reward_clip)
        
        return rewards, kl.mean().item()
    
    def compute_advantages(self, rewards, values, gamma=0.99, lam=0.95):
        """使用GAE计算优势估计"""
        advantages = torch.zeros_like(rewards)
        last_advantage = 0
        
        for t in reversed(range(len(rewards))):
            if t == len(rewards) - 1:
                next_value = 0
            else:
                next_value = values[t + 1]
            
            delta = rewards[t] + gamma * next_value - values[t]
            advantages[t] = last_advantage = delta + gamma * lam * last_advantage
        
        # 标准化优势
        advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
        return advantages
    
    def ppo_step(self, prompts, prompt_lengths):
        """一个PPO训练步"""
        batch_size = len(prompts)
        
        # 1. 生成回复
        responses = self.generate_responses(prompts, max_length=32)
        
        # 2. 计算旧策略的对数概率
        with torch.no_grad():
            old_log_probs = self.get_log_probs(self.policy, responses, responses)
        
        # 3. 计算奖励
        rewards, kl_div = self.compute_rewards(responses, prompt_lengths)
        
        # 4. 计算价值估计
        with torch.no_grad():
            logits, hidden = self.policy(responses, return_hidden=True)
            values = self.value_head(hidden[:, -1, :]).squeeze(-1)
        
        # 5. 计算优势
        advantages = self.compute_advantages(rewards, values)
        
        # 6. PPO多轮更新
        metrics = {'loss': 0, 'kl': kl_div, 'reward': rewards.mean().item()}
        
        for epoch in range(self.config.ppo_epochs):
            # 计算新策略的对数概率
            new_log_probs = self.get_log_probs(self.policy, responses, responses)
            
            # 概率比
            ratio = torch.exp(new_log_probs - old_log_probs)
            
            # PPO clip损失
            surr1 = ratio * advantages
            surr2 = torch.clamp(
                ratio, 1 - self.config.ppo_clip, 1 + self.config.ppo_clip
            ) * advantages
            policy_loss = -torch.min(surr1, surr2).mean()
            
            # 价值函数损失
            logits, hidden = self.policy(responses, return_hidden=True)
            new_values = self.value_head(hidden[:, -1, :]).squeeze(-1)
            value_loss = F.mse_loss(new_values, rewards)
            
            # 总损失
            total_loss = policy_loss + 0.5 * value_loss
            
            # 反向传播
            self.optimizer.zero_grad()
            self.value_optimizer.zero_grad()
            total_loss.backward()
            torch.nn.utils.clip_grad_norm_(self.policy.parameters(), 0.5)
            self.optimizer.step()
            self.value_optimizer.step()
            
            metrics['loss'] += total_loss.item()
        
        metrics['loss'] /= self.config.ppo_epochs
        return metrics
    
    def train(self, prompts_dataset, n_iterations=100):
        """PPO训练循环"""
        print("开始PPO训练...")
        
        for iteration in range(n_iterations):
            # 采样一批prompt
            batch = random.sample(prompts_dataset, min(self.config.batch_size, len(prompts_dataset)))
            prompts = [torch.tensor(p, device=self.device, dtype=torch.long) for p in batch]
            prompt_lengths = [len(p) for p in batch]
            
            # 标准化长度
            max_prompt_len = max(len(p) for p in prompts)
            padded_prompts = []
            for p in prompts:
                if len(p) < max_prompt_len:
                    p = F.pad(p, (0, max_prompt_len - len(p)), value=0)
                padded_prompts.append(p)
            
            metrics = self.ppo_step(padded_prompts, prompt_lengths)
            
            if iteration % 10 == 0:
                print(f"Iteration {iteration}: "
                      f"Loss={metrics['loss']:.4f}, "
                      f"KL={metrics['kl']:.4f}, "
                      f"Reward={metrics['reward']:.4f}")


def test_ppo():
    """测试PPO训练"""
    config = RLHFConfig(dim=128, n_layers=2, n_heads=4, batch_size=4, 
                        ppo_epochs=2, max_response_length=64)
    
    # 创建模型
    policy = SimpleLanguageModel(vocab_size=1000, dim=128, n_layers=2, n_heads=4)
    reference = SimpleLanguageModel(vocab_size=1000, dim=128, n_layers=2, n_heads=4)
    reference.load_state_dict(policy.state_dict())  # 参考模型初始化为策略模型
    
    base_for_rm = SimpleLanguageModel(vocab_size=1000, dim=128, n_layers=2, n_heads=4)
    reward_model = RewardModel(base_for_rm)
    
    trainer = PPOTrainer(policy, reference, reward_model, config)
    
    # 模拟prompt数据
    prompts = [[torch.randint(1, 1000, (10,)).tolist()] for _ in range(20)]
    prompts = [p[0] for p in prompts]
    
    trainer.train(prompts, n_iterations=20)

if __name__ == "__main__":
    test_ppo()

六、DPO:RLHF的简化替代方案

DPO(Direct Preference Optimization)直接使用偏好数据优化策略模型,跳过了奖励模型训练和PPO强化学习的复杂流程。DPO的理论基础是:最优的奖励函数可以用策略模型的概率比来表示,因此可以直接从偏好数据中推导出最优策略,无需显式的奖励模型。

DPO的损失函数为:

LDPO=logσ(βlogπθ(ywx)πref(ywx)βlogπθ(ylx)πref(ylx))\mathcal{L}_{DPO} = -\log\sigma\left(\beta\log\frac{\pi_\theta(y_w|x)}{\pi_{ref}(y_w|x)} - \beta\log\frac{\pi_\theta(y_l|x)}{\pi_{ref}(y_l|x)}\right)

其中 β\beta 是温度参数,πθ\pi_\theta 是策略模型,πref\pi_{ref} 是参考模型。

class DPOTrainer:
    """DPO训练器"""
    
    def __init__(
        self,
        policy_model: SimpleLanguageModel,
        reference_model: SimpleLanguageModel,
        config: RLHFConfig,
        beta: float = 0.1
    ):
        self.policy = policy_model
        self.reference = reference_model
        self.config = config
        self.beta = beta
        self.device = config.device
        
        self.policy.to(self.device)
        self.reference.to(self.device)
        
        # 冻结参考模型
        for param in self.reference.parameters():
            param.requires_grad = False
        
        self.optimizer = torch.optim.AdamW(
            policy_model.parameters(),
            lr=1e-5,
            weight_decay=0.0,
            betas=(0.9, 0.999)
        )
    
    def get_sequence_log_prob(self, model, input_ids):
        """计算序列的对数概率"""
        logits = model(input_ids)
        log_probs = F.log_softmax(logits, dim=-1)
        
        # 对齐并收集
        log_probs_labels = log_probs[:, :-1, :].gather(
            -1, input_ids[:, 1:].unsqueeze(-1)
        ).squeeze(-1)
        
        # 对每个序列求和
        sequence_log_prob = log_probs_labels.sum(dim=-1)
        return sequence_log_prob
    
    def compute_dpo_loss(
        self,
        chosen_ids: torch.Tensor,
        rejected_ids: torch.Tensor
    ) -> Tuple[torch.Tensor, dict]:
        """计算DPO损失"""
        # 策略模型的对数概率
        policy_chosen_logp = self.get_sequence_log_prob(self.policy, chosen_ids)
        policy_rejected_logp = self.get_sequence_log_prob(self.policy, rejected_ids)
        
        # 参考模型的对数概率
        with torch.no_grad():
            ref_chosen_logp = self.get_sequence_log_prob(self.reference, chosen_ids)
            ref_rejected_logp = self.get_sequence_log_prob(self.reference, rejected_ids)
        
        # 计算log-ratios
        chosen_logratios = policy_chosen_logp - ref_chosen_logp
        rejected_logratios = policy_rejected_logp - ref_rejected_logp
        
        # DPO损失
        logits = self.beta * (chosen_logratios - rejected_logratios)
        loss = -F.logsigmoid(logits).mean()
        
        # 计算指标
        with torch.no_grad():
            chosen_rewards = self.beta * chosen_logratios
            rejected_rewards = self.beta * rejected_logratios
            accuracy = (chosen_rewards > rejected_rewards).float().mean()
            margin = (chosen_rewards - rejected_rewards).mean()
        
        return loss, {
            'loss': loss.item(),
            'accuracy': accuracy.item(),
            'margin': margin.item(),
            'chosen_reward': chosen_rewards.mean().item(),
            'rejected_reward': rejected_rewards.mean().item()
        }
    
    def train_step(self, chosen_ids, rejected_ids):
        """训练一步"""
        self.policy.train()
        
        loss, metrics = self.compute_dpo_loss(chosen_ids, rejected_ids)
        
        self.optimizer.zero_grad()
        loss.backward()
        torch.nn.utils.clip_grad_norm_(self.policy.parameters(), 1.0)
        self.optimizer.step()
        
        return metrics
    
    def train(self, examples: List[PreferenceExample], tokenizer, epochs=2):
        """完整DPO训练"""
        batch_size = self.config.batch_size
        
        for epoch in range(epochs):
            random.shuffle(examples)
            total_metrics = {}
            n_batches = 0
            
            for i in range(0, len(examples), batch_size):
                batch = examples[i:i+batch_size]
                
                # 准备数据
                chosen_texts = [ex.prompt + ex.response_chosen for ex in batch]
                rejected_texts = [ex.prompt + ex.response_rejected for ex in batch]
                
                chosen_ids = [tokenizer.encode(t) for t in chosen_texts]
                rejected_ids = [tokenizer.encode(t) for t in rejected_texts]
                
                # 填充
                max_len = max(max(len(c) for c in chosen_ids),
                             max(len(r) for r in rejected_ids))
                max_len = min(max_len, 64)
                
                chosen_tensor = torch.tensor([
                    c[:max_len] + [0] * (max_len - len(c[:max_len]))
                    for c in chosen_ids
                ], device=self.device)
                rejected_tensor = torch.tensor([
                    r[:max_len] + [0] * (max_len - len(r[:max_len]))
                    for r in rejected_ids
                ], device=self.device)
                
                metrics = self.train_step(chosen_tensor, rejected_tensor)
                
                for k, v in metrics.items():
                    total_metrics[k] = total_metrics.get(k, 0) + v
                n_batches += 1
                
                if n_batches % 5 == 0:
                    print(f"Epoch {epoch+1}, Batch {n_batches}: "
                          f"Loss={metrics['loss']:.4f}, "
                          f"Acc={metrics['accuracy']:.2%}")
            
            avg_metrics = {k: v / n_batches for k, v in total_metrics.items()}
            print(f"Epoch {epoch+1} 完成: Loss={avg_metrics['loss']:.4f}, "
                  f"Acc={avg_metrics['accuracy']:.2%}")


def test_dpo():
    """测试DPO训练"""
    config = RLHFConfig(dim=128, n_layers=2, n_heads=4, batch_size=4)
    
    policy = SimpleLanguageModel(vocab_size=1000, dim=128, n_layers=2, n_heads=4)
    reference = SimpleLanguageModel(vocab_size=1000, dim=128, n_layers=2, n_heads=4)
    reference.load_state_dict(policy.state_dict())
    
    trainer = DPOTrainer(policy, reference, config, beta=0.1)
    tokenizer = SimpleTokenizer(vocab_size=1000)
    
    # 模拟偏好数据
    examples = [
        PreferenceExample(
            prompt="What is AI?",
            response_chosen="AI is artificial intelligence, a branch of computer science.",
            response_rejected="AI is bad and useless."
        ),
        PreferenceExample(
            prompt="Explain ML.",
            response_chosen="Machine learning enables systems to learn from data.",
            response_rejected="ML is hard."
        ),
    ] * 10
    
    trainer.train(examples, tokenizer, epochs=3)

if __name__ == "__main__":
    test_dpo()

七、RLHF的工程挑战与实践建议

RLHF在实际工程中面临多项挑战。首先是偏好数据质量:标注者需要专业知识且一致性要求高,数据采集成本高。其次是奖励模型的泛化性:分布外输入可能导致奖励值不准。第三是PPO训练的不稳定性:超参数敏感,容易出现策略崩溃或KL散度爆炸。第四是计算成本:PPO需要同时维护策略模型、参考模型、奖励模型和价值模型,显存压力大。

实践建议包括:SFT阶段使用高质量的指令数据,数据量不需要太大但质量要高;奖励模型训练时确保训练集和评估集的标注一致性;PPO训练中KL系数从0开始逐步增大,防止初期不稳定;考虑使用DPO替代PPO以简化流程,DPO在多数场景下能达到接近PPO的效果但训练更简单稳定。

def rlhf_comparison():
    """RLHF方法对比"""
    methods = [
        {
            "name": "RLHF (SFT + RM + PPO)",
            "stages": "3阶段",
            "reward_model": "需要",
            "online_sampling": "需要",
            "stability": "较低",
            "compute": "高",
            "quality": "最佳(数据充足时)",
            "complexity": "高"
        },
        {
            "name": "DPO",
            "stages": "2阶段(SFT+DPO)",
            "reward_model": "不需要",
            "online_sampling": "不需要",
            "stability": "较高",
            "compute": "中",
            "quality": "接近RLHF",
            "complexity": "中"
        },
        {
            "name": "ORPO",
            "stages": "1阶段",
            "reward_model": "不需要",
            "online_sampling": "不需要",
            "stability": "高",
            "compute": "低",
            "quality": "良好",
            "complexity": "低"
        },
        {
            "name": "KTO (Kahneman-Tversky)",
            "stages": "2阶段",
            "reward_model": "不需要",
            "online_sampling": "不需要",
            "stability": "高",
            "compute": "低",
            "quality": "良好",
            "complexity": "低"
        },
    ]
    
    print(f"{'方法':<25} {'阶段':<15} {'RM':<8} {'采样':<8} {'稳定':<6} {'算力':<6} {'质量':<20}")
    print("-" * 95)
    for m in methods:
        print(f"{m['name']:<25} {m['stages']:<15} {m['reward_model']:<8} "
              f"{m['online_sampling']:<8} {m['stability']:<6} {m['compute']:<6} "
              f"{m['quality']:<20}")
    
    print("\n选择建议:")
    print("  - 有充足算力和偏好数据: RLHF (PPO)")
    print("  - 想要效果好但简化流程: DPO")
    print("  - 资源有限追求效率: ORPO/KTO")

if __name__ == "__main__":
    rlhf_comparison()

八、总结

RLHF通过人类偏好数据引导模型对齐,是现代大模型从"能生成文本"到"能遵循人类意图"的关键技术桥梁。三阶段流程中,SFT建立基本能力,RM学习偏好信号,PPO利用奖励信号优化策略。DPO作为简化替代方案,跳过奖励模型和强化学习,直接从偏好数据优化策略,训练更简单且稳定性更高。实际工程中需要根据算力、数据质量和任务需求选择合适的对齐方法。随着对齐技术的持续演进(如宪法AI、自我对弈等),大模型的安全性和有用性将不断提升,推动AI向更负责任的方向发展。

【声明】本内容来自华为云开发者社区博主,不代表华为云及华为云开发者社区的观点和立场。转载时必须标注文章的来源(华为云社区)、文章链接、文章作者等基本信息,否则作者和本社区有权追究责任。如果您发现本社区中有涉嫌抄袭的内容,欢迎发送邮件进行举报,并提供相关证据,一经查实,本社区将立刻删除涉嫌侵权内容,举报邮箱: cloudbbs@huaweicloud.com
  • 点赞
  • 收藏
  • 关注作者

评论(0

0/1000
抱歉,系统识别当前为高风险访问,暂不支持该操作

全部回复

上滑加载中

设置昵称

在此一键设置昵称,即可参与社区互动!

*长度不超过10个汉字或20个英文字符,设置后3个月内不可修改。

*长度不超过10个汉字或20个英文字符,设置后3个月内不可修改。