跳转至

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 架构的训练侧基石。