跳转至

dp_actor.py — FSDP 数据并行 Actor

文件概述

基于 FSDP 的 DataParallelPPOActor 实现(约 677 行)。是老架构(fsdp_workers.py)中 Actor 的核心实现,包含模型前向计算、log_prob 计算、PPO loss 计算和策略更新的完整逻辑。

核心类

DataParallelPPOActor

class DataParallelPPOActor(BasePPOActor):
    def __init__(self, config, actor_module, actor_optimizer, ...):
        self.actor_module = actor_module  # FSDP 包装的模型
        self.actor_optimizer = actor_optimizer
        self.use_remove_padding = config.use_remove_padding  # 是否去除 padding

核心方法

1. _forward_micro_batch - 微批次前向

def _forward_micro_batch(self, micro_batch):
    """对一个微批次执行前向传播,计算 log_prob 和 entropy"""
    input_ids = micro_batch['input_ids']
    attention_mask = micro_batch['attention_mask']
    position_ids = micro_batch['position_ids']

    # 调用 HuggingFace 模型
    output = self.actor_module(
        input_ids=input_ids,
        attention_mask=attention_mask,
        position_ids=position_ids,
    )
    logits = output.logits

    # 从 logits 计算 log_prob
    log_prob = logprobs_from_logits(logits, labels)
    return log_prob, entropy

2. compute_log_prob - 计算 log 概率

def compute_log_prob(self, data: DataProto) -> DataProto:
    """计算每个 response token 的 log 概率

    支持特性:
    - Remove Padding: 去除无效 padding,减少计算量
    - Ulysses SP: 序列并行,处理超长序列
    - 微批次处理: 大 batch 分成小批次避免 OOM
    """
    # 1. 数据预处理(去 padding、序列并行分片)
    # 2. 分微批次前向计算
    # 3. 收集结果
    ...

3. update_policy - PPO 策略更新

def update_policy(self, data: DataProto) -> dict:
    """执行多轮 PPO 更新

    流程:
    1. 遍历 ppo_epochs 轮
    2. 每轮将数据分成 mini-batch
    3. 每个 mini-batch 做前向+反向+优化器步骤
    4. 计算 PPO clip loss + entropy loss + KL loss
    """
    for epoch in range(ppo_epochs):
        for mini_batch in data.split(mini_batch_size):
            # 前向计算 log_prob
            log_prob = self._forward_micro_batch(mini_batch)

            # 计算 PPO loss
            ratio = (log_prob - old_log_prob).exp()
            surr1 = ratio * advantages
            surr2 = ratio.clamp(1-eps, 1+eps) * advantages
            pg_loss = -torch.min(surr1, surr2).mean()

            # 反向传播 + 优化器更新
            pg_loss.backward()
            optimizer.step()

重要特性

Remove Padding 优化

# 传统方式: 填充到相同长度(浪费计算)
# [hello, world, <pad>, <pad>]  shape: (batch, max_len)
# [hi, there, now, <pad>]

# Remove Padding: 拼接有效 token(高效)
# [hello, world, hi, there, now]  shape: (total_tokens,)
# 配合 cu_seqlens 记录每个序列的边界

动态批大小

通过 max_token_len 参数,根据序列实际长度动态调整微批次大小,确保每个微批次的总 token 数相近,提高 GPU 利用率。

与其他模块的关系

  • 继承自 base.py 的 BasePPOActor
  • 被 fsdp_workers.py 创建和使用
  • 使用 sharding_manager/fsdp_ulysses.py 处理序列并行数据分片
  • 损失计算逻辑与 utils/losses.py 中的 ppo_loss 类似

小结

DataParallelPPOActor 是 FSDP 后端下 Actor 的完整实现,支持 Remove Padding、Ulysses 序列并行、动态批大小等优化。对于理解 PPO 训练的前向和更新流程非常有价值。