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 训练的前向和更新流程非常有价值。