transformer_impl.py — FSDP 训练引擎实现¶
文件概述¶
基于 PyTorch FSDP 的训练引擎实现(约 1000+ 行),是最常用的训练后端。支持 FSDP 分片、Ulysses 序列并行、LoRA、Remove Padding、混合精度等特性。
核心类¶
FSDPEngine(BaseEngine)¶
FSDP 引擎基类,实现了所有 BaseEngine 的抽象方法。
class FSDPEngine(BaseEngine):
def __init__(self, model_config, engine_config, optimizer_config, checkpoint_config):
# 初始化设备网格
self.device_mesh = create_device_mesh(world_size, fsdp_size)
# Ulysses 序列并行
if ulysses_sp_size > 1:
self.ulysses_device_mesh = init_device_mesh(...)
initialize - 初始化¶
def initialize(self):
"""构建模型 → FSDP 包装 → 优化器 → 学习率调度器"""
self._build_model_optimizer()
self.checkpoint_manager = FSDPCheckpointManager(...)
# 根据卸载策略将参数/优化器移到 CPU
self.to(device="cpu", model=offload_param, optimizer=offload_optimizer)
_build_module - 构建模型¶
def _build_module(self):
"""从 HuggingFace 加载模型"""
module = AutoModelForCausalLM.from_pretrained(...)
# 可选: 应用 Liger 内核、融合内核、monkey patch
apply_monkey_patch(module, use_remove_padding=True, ...)
# 可选: 梯度检查点
module.gradient_checkpointing_enable(...)
return module
forward_backward_batch - 前向+反向¶
def forward_backward_batch(self, data, loss_function, forward_only=False):
"""分微批次执行前向和反向传播"""
micro_batches, indices = prepare_micro_batches(data, ...)
for micro_batch in micro_batches:
with ctx: # torch.no_grad() if forward_only
loss, output = self.forward_step(micro_batch, loss_function, forward_only)
if not forward_only:
loss.backward()
return postprocess_batch_func(output_lst, indices, data)
to - 设备迁移¶
def to(self, device, model=True, optimizer=True, grad=True):
"""将模型/优化器在 GPU 和 CPU 之间迁移
GPU → CPU: 节省显存(Hybrid Engine 训练阶段结束后)
CPU → GPU: 恢复训练(Hybrid Engine 训练阶段开始时)
"""
FSDPEngineWithLMHead - 语言模型引擎¶
@EngineRegistry.register(model_type="language_model", backend=["fsdp"], device=["cuda", "npu"])
class FSDPEngineWithLMHead(FSDPEngine):
"""带 LM Head 的 FSDP 引擎"""
def forward_step(self, micro_batch, loss_function, forward_only):
"""一个微批次的完整前向步骤"""
# 1. 准备输入
input_ids, position_ids, attention_mask = ...
# 2. 模型前向
output = self.module(input_ids=input_ids, ...)
logits = output.logits
# 3. 计算 log_prob
logits /= temperature
log_prob = logprobs_from_logits(logits, labels)
# 4. 可选: 计算 entropy
if calculate_entropy:
entropy = entropy_from_logits(logits)
# 5. 计算损失
if loss_function:
loss, metrics = loss_function(model_output, data)
return loss, output_dict
FSDPEngineWithValueHead - 价值模型引擎¶
@EngineRegistry.register(model_type="value_model", backend=["fsdp"], device=["cuda", "npu"])
class FSDPEngineWithValueHead(FSDPEngine):
"""Critic 价值模型的 FSDP 引擎"""
# 输出是标量 value 而非 logits
get_per_tensor_param - 权重同步¶
def get_per_tensor_param(self, base_sync_done=False, **kwargs):
"""获取每层参数用于同步到推理引擎
支持:
- 全量权重同步
- LoRA adapter 单独同步
- FSDP 分片权重的全聚合
"""
load_fsdp_model_to_gpu(self.module) # 确保在 GPU 上
params = self.module.state_dict() # 获取完整 state_dict
params = convert_weight_keys(params) # 转换键名格式
return params, peft_config
上下文管理器¶
EngineTrainModeCtx¶
class EngineTrainModeCtx(BaseEngineCtx):
def __enter__(self):
super().__enter__() # 加载模型/优化器到 GPU
self.engine.module.train() # 切换到训练模式
def __exit__(self, ...):
self.engine.optimizer_zero_grad() # 清零梯度
super().__exit__(...) # 卸载到 CPU
EngineEvalModeCtx¶
class EngineEvalModeCtx(BaseEngineCtx):
def __enter__(self):
super().__enter__() # 加载模型到 GPU(不加载优化器)
self.engine.module.eval() # 切换到推理模式
与其他模块的关系¶
- 继承
engine/base.py的BaseEngine - 通过
EngineRegistry注册 - 被
engine_workers.py的TrainingWorker使用 - 使用
engine/fsdp/utils.py的设备网格和分片策略工具 - 使用
utils/losses.py中的损失函数
小结¶
FSDPEngine 是最常用的训练引擎,通过 FSDP 将大模型分片到多个 GPU,支持参数/优化器 CPU 卸载,是 Hybrid Engine 架构的训练侧基石。