跳转至

transformer_impl.py — TorchTitan 训练引擎实现

文件概述

基于 PyTorch TorchTitan 的训练引擎实现(约 734 行),使用 FSDP2 + TP + PP + EP 并行策略。TorchTitan 是 PyTorch 官方的分布式训练框架。

核心类

TorchTitanEngine(BaseEngine)

class TorchTitanEngine(BaseEngine):
    def __init__(self, model_config, engine_config, ...):
        # 从 HF config 自动推导 torchtitan 模型名和变体
        name, flavor = derive_torchtitan_name_and_flavor(self.model_config.hf_config)

        # 获取 ModelSpec
        model_module = importlib.import_module(f"torchtitan.models.{name}")
        model_spec = model_module.model_registry(flavor)

        # 构建 TorchTitan Trainer
        self.config = Trainer.Config(
            model_spec=model_spec,
            optimizer=OptimizersContainer.Config(...),
            lr_scheduler=LRSchedulersContainer.Config(...),
            parallelism=ParallelismConfig(...),
            ...
        )
        self.trainer = Trainer(self.config)

并行维度配置

parallelism = ParallelismConfig(
    data_parallel_shard_degree=...,       # FSDP2 分片维度
    data_parallel_replicate_degree=...,   # DDP 复制维度
    tensor_parallel_degree=...,           # 张量并行
    pipeline_parallel_degree=...,         # 流水线并行
    context_parallel_degree=...,          # 上下文并行
    expert_parallel_degree=...,           # 专家并行
)

初始化

def initialize(self):
    self.module = self.trainer.model_parts
    self.checkpointer = self.trainer.checkpointer
    self.checkpointer.load()  # 加载初始 HF 权重
    self.optimizer = self.trainer.optimizers
    self.lr_scheduler = self.trainer.lr_schedulers

TorchTitanEngineWithLMHead

@EngineRegistry.register(model_type="language_model", backend=["torchtitan"], device=["cuda", "npu"])
class TorchTitanEngineWithLMHead(TorchTitanEngine):
    def forward_step(self, micro_batch, loss_function, forward_only):
        """使用 torch.autocast 和 TorchTitan 的训练上下文"""
        with torch.autocast(device_type=device_name, dtype=torch.bfloat16):
            logits = self.model_forward_step(inputs=input_ids, ...)
            model_output = self.prepare_model_outputs(logits, ...)
            loss, metrics = loss_function(model_output, micro_batch)
        return loss, output

权重同步(EP 支持)

def get_per_tensor_param(self, **kwargs):
    """获取权重,支持 Expert Parallel 的 all-gather"""
    params = get_model_state_dict(module)

    # HuggingFace 键名转换
    params = sd_adapter.to_hf(params)

    if self.parallel_dims.ep_enabled:
        # EP 模式: 需要跨 EP rank all-gather 专家权重
        per_tensor_param = iter_per_tensor_params_ep(params, device, ep_group, ep_size)
    else:
        # 普通模式: DTensor → full_tensor
        per_tensor_param = ((name, param.full_tensor()) for name, param in params.items())

与 FSDP Engine 的区别

特性 FSDPEngine TorchTitanEngine
FSDP 版本 FSDP1 FSDP2
TP 支持 无 有
PP 支持 无 有(开发中)
EP 支持 无 有
CP 支持 无 有
模型来源 HuggingFace TorchTitan ModelSpec

与其他模块的关系

  • 继承 engine/base.py 的 BaseEngine
  • 使用 engine/torchtitan/utils.py 的辅助工具
  • 使用 TorchTitan 框架的 Trainer、Checkpoint 等组件

小结

TorchTitan 引擎利用 PyTorch 原生的分布式能力,提供了 FSDP2 + TP + EP + CP 的现代化并行组合。