跳转至

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 自动管理资源的分配和释放。