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 的现代化并行组合。