已开启
【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 5 Trainer 接缝层 #143
changzherui创建于 5月15日
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M5 Trainer 接缝层”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M5”
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M5 Trainer 接缝层”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M5”
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M5 Trainer 接缝层”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M5 Trainer 接缝层”
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M5 Trainer 接缝层”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M5 Trainer 接缝层”
5月25日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 5 Trainer 接缝层”,原标题为“【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M5 Trainer 接缝层”
5月25日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 5 Trainer 接缝层”,原标题为“【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M5 Trainer 接缝层”
5月25日 修改了issue 的描述
Part 5 —
trainer/接缝层1. 目标
把 M2 + M3 + M1 拼成完整的"配置 → 训练器"链路:
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 微调)ModelSpecv2 字段:name / flavor / model: BaseModel.Config | None / build_model_fn / parallelize_fn / pipelining_fn / post_optimizer_build_fn / state_dict_adapter / clip_grad_fnhyper_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拆分要点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)_build_optimizer(base.py:555-604)isinstance(self.args, HyperTrainer.Config)→ 调self.args.optimizer.build(...);否则走旧逻辑旧
LLMTrainer(args)路径保持完全不变(llm_trainer.py:47-66不动)。6. 与 torchtitan 接口差异说明
parallelize_fn签名扩展:torchtitan 是(model, parallel_dims, *, training, parallelism, ...),hyper 需兼容旧(model, mesh, cfg)(models/qwen3_5/parallelize.py:102)ModelSpec在 M1 增加新字段;新 spec 用新签名,旧 spec 仍可用 v1 签名HyperTrainer.Config.model_spec用Annotated[..., tyro.conf.Suppress]从 CLI 排除ModelSpec持 callable / dataclass,tyro 不能解析commconfig 单独抽出,旧值由_legacy_to_trainer_config写入train.comm_backend散着;新版分组更清晰base.py:680)MetricsProcessor只是组件式入口,M8 再统一7. 开发步骤
BaseTrainer.__init__中可复用的初始化逻辑抽 helper(_helper_setup_distributed / _helper_post_parallelize / _helper_load_weights)。旧调用点保持调 helper 包装方法,行为 0 变化。trainer_config.py的HyperTrainer.Config,引用 M2 / M3 的所有 Config。hyper_trainer.py的HyperTrainer.__init__11 步。protocols/model_spec.py到 v2 字段集(M1 已基本就位,M5 微调和测试)。8. 验证标准
新建
tests/torch/ut/trainer/:test_trainer_config_build.pyBaseModel.Config+ mock 所有组件,构造HyperTrainer.Config(...).build(),断言每个组件build被调一次且参数正确(按 §4 顺序)test_trainer_branch.pyLLMTrainer(parse_args(HyperTrainerConfig)),②HyperTrainer(_legacy_to_trainer_config(parse_args(HyperTrainerConfig))),断言self.model / self.optimizer / self.lr_scheduler结构等价(pname 集合、shape 一致)test_base_trainer_unchanged.pytests/torch/st/qwen3_5/*)跑通,结果与 main 分支 bit-exact通过门槛:
LLMTrainer(args)路径 0 行为改动。scripts/train_lm.py不动(M7 才动)。9. 工期 & 依赖