base.py — 训练引擎基类与注册表¶
文件概述¶
定义训练引擎的抽象基类 BaseEngine、上下文管理器 BaseEngineCtx 和引擎注册表 EngineRegistry。这是统一引擎架构的核心文件。
核心类¶
BaseEngine(ABC) - 引擎抽象基类¶
class BaseEngine(ABC):
"""所有训练引擎必须实现的接口"""
@abstractmethod
def initialize(self):
"""初始化模型、优化器、学习率调度器"""
pass
@abstractmethod
def train_mode(self) -> ContextManager:
"""返回训练模式上下文管理器"""
pass
@abstractmethod
def eval_mode(self) -> ContextManager:
"""返回推理模式上下文管理器"""
pass
@abstractmethod
def forward_backward_batch(self, data, loss_function, forward_only=False):
"""执行前向(+反向)传播"""
pass
@abstractmethod
def optimizer_step(self) -> float:
"""优化器更新,返回 grad_norm"""
pass
@abstractmethod
def lr_scheduler_step(self) -> float:
"""学习率调度器步进,返回当前 lr"""
pass
@abstractmethod
def to(self, device, model=True, optimizer=True, grad=True):
"""将模型/优化器移到指定设备(GPU/CPU 卸载)"""
pass
@abstractmethod
def get_per_tensor_param(self, **kwargs):
"""获取每层参数用于权重同步到推理引擎"""
pass
@abstractmethod
def save_checkpoint(self, local_path, ...):
"""保存检查点"""
pass
@abstractmethod
def load_checkpoint(self, local_path, ...):
"""加载检查点"""
pass
BaseEngineCtx - 引擎上下文管理器¶
class BaseEngineCtx:
"""管理 train/eval 模式切换,自动处理参数/优化器的 GPU/CPU 迁移"""
def __init__(self, engine, mode="train"):
self.engine = engine
self.mode = mode
def __enter__(self):
# 训练模式: 加载模型和优化器到 GPU
if self.mode == "train":
self.engine.to(device="cuda", model=True, optimizer=True)
# 推理模式: 只加载模型到 GPU
else:
self.engine.to(device="cuda", model=True, optimizer=False)
def __exit__(self, ...):
# 根据卸载策略决定是否卸载到 CPU
if self.engine.is_param_offload_enabled:
self.engine.to(device="cpu", model=True)
if self.engine.is_optimizer_offload_enabled:
self.engine.to(device="cpu", optimizer=True)
EngineRegistry - 引擎注册表¶
class EngineRegistry:
"""通过装饰器注册引擎,通过 (model_type, backend, device) 查找"""
_registry = {}
@classmethod
def register(cls, model_type, backend, device="cuda"):
"""装饰器:注册引擎类"""
def decorator(engine_cls):
cls._registry[(model_type, backend, device)] = engine_cls
return engine_cls
return decorator
@classmethod
def get(cls, model_type, backend, device="cuda"):
"""查找注册的引擎类"""
return cls._registry[(model_type, backend, device)]
使用示例:
# 注册
@EngineRegistry.register(model_type="language_model", backend="fsdp")
class FSDPEngineWithLMHead(FSDPEngine):
...
# 查找
engine_cls = EngineRegistry.get(model_type="language_model", backend="fsdp")
engine = engine_cls(model_config, engine_config, ...)
设计思想¶
EngineRegistry
├── ("language_model", "fsdp", "cuda") → FSDPEngineWithLMHead
├── ("value_model", "fsdp", "cuda") → FSDPEngineWithValueHead
├── ("language_model", "megatron", "cuda") → MegatronEngineWithLMHead
├── ("language_model", "torchtitan", "cuda")→ TorchTitanEngineWithLMHead
├── ("language_model", "veomni", "cuda") → VeOmniEngineWithLMHead
└── ("language_model", "megatron", "npu") → MindspeedEngineWithLMHead
与其他模块的关系¶
- 被
engine_workers.py的TrainingWorker使用来创建引擎 - 被
engine/fsdp/,engine/megatron/等子模块继承实现 BaseEngineCtx自动管理训练/推理模式切换和 CPU 卸载
小结¶
base.py 是统一引擎架构的基石。BaseEngine 定义了所有训练引擎的统一接口,EngineRegistry 实现了后端可插拔的注册机制,BaseEngineCtx 自动管理资源的分配和释放。