已开启
【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 5 Trainer 接缝层 #143
changzherui创建于  5月15日
changzherui
changzherui成员
5月15日 创建

Part 5 — trainer/ 接缝层

在不动 LLMTrainer / VLTrainer 的前提下,给 BaseTrainer 加一个"新 Config 入口"分支,以及 HyperTrainer.Config 顶层 Configurable 树。新 / 旧两条路径共享同一份算法实现(通过抽公共 helper)。


1. 目标

把 M2 + M3 + M1 拼成完整的"配置 → 训练器"链路:

HyperTrainer.Config(Configurable 树)
   .build()
      └─ HyperTrainer.__init__(config)
            ├─ init_distributed + ParallelDims
            ├─ tokenizer / dataloader / model / parallelize / optimizer / lr / ckpt / metrics / profiler
            └─ trainer.train()

2. 任务边界

新文件 / 修改文件 内容
hyper_parallel/trainer/trainer_config.py(新) class HyperTrainer.Config(Configurable.Config):顶层 Configurable 树(见 §3)
hyper_parallel/trainer/hyper_trainer.py(新) class HyperTrainer(BaseTrainer):新 Config 入口的 trainer 子类。__init__(self, config: HyperTrainer.Config) 调一连串 config.<comp>.build()
hyper_parallel/protocols/model_spec.py(M1 已落,M5 微调) 验收时确认 ModelSpec v2 字段:name / flavor / model: BaseModel.Config | None / build_model_fn / parallelize_fn / pipelining_fn / post_optimizer_build_fn / state_dict_adapter / clip_grad_fn
hyper_parallel/trainer/base.py(改) 抽出 _helper_setup_distributed / _helper_post_parallelize / ... 等纯函数 helper(不改可见行为)。BaseTrainer.__init__ 0 改动;helper 同时被旧 _build_* 和新 HyperTrainer.__init__ 调用

3. HyperTrainer.Config 结构

from dataclasses import dataclass, field
from typing import Annotated

import tyro

from hyper_parallel.protocols import Configurable, ModelSpec
from hyper_parallel.config import (
    TrainingConfig, ParallelismConfig, ActivationCheckpointConfig,
    CompileConfig, CommConfig, DebugConfig,
)
from hyper_parallel.components.optimizer import OptimizersContainer
from hyper_parallel.components.lr_scheduler import LRSchedulersContainer
from hyper_parallel.components.loss import CrossEntropyLoss, BaseLoss
from hyper_parallel.components.tokenizer import HuggingFaceTokenizer, BaseTokenizer
from hyper_parallel.components.dataloader import BaseDataLoader
from hyper_parallel.components.checkpoint import CheckpointManager
from hyper_parallel.components.profiler import Profiler
from hyper_parallel.components.metrics import MetricsProcessor


class HyperTrainer(Configurable):

    @dataclass(kw_only=True, slots=True)
    class Config(Configurable.Config):
        # ModelSpec 持 callable / dataclass,tyro 无法解析,对 CLI 隐藏
        model_spec: Annotated[ModelSpec | None, tyro.conf.Suppress] = None

        hf_assets_path: str = "outputs/hf_assets"
        dump_folder: str = "outputs"

        tokenizer:             BaseTokenizer.Config         = field(default_factory=HuggingFaceTokenizer.Config)
        dataloader:            BaseDataLoader.Config        = field(default_factory=BaseDataLoader.Config)
        optimizer:             OptimizersContainer.Config   = field(default_factory=OptimizersContainer.Config)
        lr_scheduler:          LRSchedulersContainer.Config = field(default_factory=LRSchedulersContainer.Config)
        loss:                  BaseLoss.Config              = field(default_factory=CrossEntropyLoss.Config)
        checkpoint:            CheckpointManager.Config     = field(default_factory=CheckpointManager.Config)
        metrics:               MetricsProcessor.Config      = field(default_factory=MetricsProcessor.Config)
        profiler:              Profiler.Config              = field(default_factory=Profiler.Config)

        training:              TrainingConfig               = field(default_factory=TrainingConfig)
        parallelism:           ParallelismConfig            = field(default_factory=ParallelismConfig)
        activation_checkpoint: ActivationCheckpointConfig   = field(default_factory=ActivationCheckpointConfig)
        compile:               CompileConfig                = field(default_factory=CompileConfig)
        comm:                  CommConfig                   = field(default_factory=CommConfig)
        debug:                 DebugConfig                  = field(default_factory=DebugConfig)

4. HyperTrainer.__init__ 11 步

对应 torchtitan Trainer.__init__:

class HyperTrainer(BaseTrainer):
    def __init__(self, config: "HyperTrainer.Config"):
        self.config = config

        # 1. init_process_group + ParallelDims(复用 helper)
        self._helper_setup_distributed(config)

        # 2. tokenizer
        self.tokenizer = config.tokenizer.build(path=config.tokenizer.path)

        # 3. dataloader
        self.dataloader = config.dataloader.build(
            dp_world_size=self.parallel_dims.dp_world_size,
            dp_rank=self._dp_rank(),
            tokenizer=self.tokenizer,
            seq_len=config.training.seq_len,
        )

        # 4. model config 接入运行时
        model_config = config.model_spec.model
        model_config.update_from_config(trainer_config=config)
        # ←—— update_from_config 内部调 set_<name>_sharding_config()

        # 5. 构造模型(meta device)
        with init_empty_weights():
            model = model_config.build()

        # 6. 协议自检
        model.verify_module_protocol()

        # 7. 并行化(CP / TP / AC / FSDP)
        model = config.model_spec.parallelize_fn(
            model, self.parallel_dims,
            training=config.training,
            parallelism=config.parallelism,
            activation_checkpoint=config.activation_checkpoint,
            compile=config.compile,
            dump_folder=config.dump_folder,
        )

        # 8. 物化 + 权重初始化
        model.to_empty(device=platform.device_type())
        model.init_weights(buffer_device=platform.device_type())
        self.model = model

        # 9. optimizer / lr_scheduler
        self.optimizer = config.optimizer.build(model_parts=[model])
        self.lr_scheduler = config.lr_scheduler.build(
            optimizers=self.optimizer,
            training_steps=config.training.steps,
        )

        # 10. checkpoint
        sd_adapter = None
        if config.model_spec.state_dict_adapter is not None:
            sd_adapter = config.model_spec.state_dict_adapter(
                model_config, config.hf_assets_path,
            )
        self.checkpointer = config.checkpoint.build(
            model_parts=[model],
            optimizers=self.optimizer,
            lr_schedulers=self.lr_scheduler,
            dataloader=self.dataloader,
            sd_adapter=sd_adapter,
        )

        # 11. metrics / profiler
        self.metrics = config.metrics.build(trainer=self)
        self.profiler = config.profiler.build()

5. BaseTrainer 拆分要点

目标:算法零分叉,新 / 旧两条 __init__ 都调用同一份 helper。

现 BaseTrainer 方法 重构后
_setup(base.py:124-192) 抽出 _helper_setup_distributed(config) 纯函数;原方法变为薄包装
_post_parallelize(base.py:425-462) 抽出 _helper_post_parallelize(model, init_device, weights_path, mp_cfg, ...)
_materialize_and_init_shards(base.py:1183-1222) 已经是 helper 风格,保持不变
_build_optimizer(base.py:555-604) 新增内部 if 分支:isinstance(self.args, HyperTrainer.Config) → 调 self.args.optimizer.build(...);否则走旧逻辑

旧 LLMTrainer(args) 路径保持完全不变(llm_trainer.py:47-66 不动)。

6. 与 torchtitan 接口差异说明

# 差异点 原因
1 parallelize_fn 签名扩展:torchtitan 是 (model, parallel_dims, *, training, parallelism, ...),hyper 需兼容旧 (model, mesh, cfg)(models/qwen3_5/parallelize.py:102) ModelSpec 在 M1 增加新字段;新 spec 用新签名,旧 spec 仍可用 v1 签名
2 HyperTrainer.Config.model_spec 用 Annotated[..., tyro.conf.Suppress] 从 CLI 排除 ModelSpec 持 callable / dataclass,tyro 不能解析
3 comm config 单独抽出,旧值由 _legacy_to_trainer_config 写入 hyper 旧字段在 train.comm_backend 散着;新版分组更清晰
4 hyper 暂时保留 13 个 Callback 体系不动(base.py:680) 减小 PR 范围;M3 MetricsProcessor 只是组件式入口,M8 再统一

7. 开发步骤

  1. Step 1(0.5 d):把 BaseTrainer.__init__ 中可复用的初始化逻辑抽 helper(_helper_setup_distributed / _helper_post_parallelize / _helper_load_weights)。旧调用点保持调 helper 包装方法,行为 0 变化。
  2. Step 2(0.5 d):写 trainer_config.py 的 HyperTrainer.Config,引用 M2 / M3 的所有 Config。
  3. Step 3(1 d):写 hyper_trainer.py 的 HyperTrainer.__init__ 11 步。
  4. Step 4(0.5 d):升级 protocols/model_spec.py 到 v2 字段集(M1 已基本就位,M5 微调和测试)。

8. 验证标准

新建 tests/torch/ut/trainer/:

测试 断言要点
test_trainer_config_build.py mock 一个 BaseModel.Config + mock 所有组件,构造 HyperTrainer.Config(...).build(),断言每个组件 build 被调一次且参数正确(按 §4 顺序)
test_trainer_branch.py 同一个 yaml 走两条路径:①LLMTrainer(parse_args(HyperTrainerConfig)),②HyperTrainer(_legacy_to_trainer_config(parse_args(HyperTrainerConfig))),断言 self.model / self.optimizer / self.lr_scheduler 结构等价(pname 集合、shape 一致)
test_base_trainer_unchanged.py 旧路径所有现存 ST 测试(tests/torch/st/qwen3_5/*)跑通,结果与 main 分支 bit-exact

通过门槛:

  • 3 个测试全绿。
  • 旧 LLMTrainer(args) 路径 0 行为改动。
  • scripts/train_lm.py 不动(M7 才动)。

9. 工期 & 依赖

工期 2.5 天
依赖 M1 + M2 + M3
不依赖 M4(用 mock 模型即可验证)
下游 M6 / M7
likedislike
changzheruichangzherui成员
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M5 Trainer 接缝层”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M5”
changzheruichangzherui成员
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M5 Trainer 接缝层”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M5 Trainer 接缝层”
changzheruichangzherui成员
5月25日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 5 Trainer 接缝层”,原标题为“【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M5 Trainer 接缝层”
changzheruichangzherui成员
5月25日 修改了issue 的描述