mindspore-ai / mindspore-ai/hyper-parallel
【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 5 Trainer 接缝层
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 53
- Forks
- 63
- Avg merge
- 23h 45m
- Merged PRs (30d)
- 63
Description
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. 开发步骤
- Step 1(0.5 d):把
BaseTrainer.__init__中可复用的初始化逻辑抽 helper(_helper_setup_distributed / _helper_post_parallelize / _helper_load_weights)。旧调用点保持调 helper 包装方法,行为 0 变化。 - Step 2(0.5 d):写
trainer_config.py的HyperTrainer.Config,引用 M2 / M3 的所有 Config。 - Step 3(1 d):写
hyper_trainer.py的HyperTrainer.__init__11 步。 - 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 |
schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 143
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/143
Contributor guide
No contributing guide indexed for this repository
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start with hyper_parallel/trainer/base.py, especially _setup, _post_parallelize, and _build_optimizer, then review the new trainer_config.py and hyper_trainer.py entry points. Implement the shared initialization path and ModelSpec v2 integration described in the issue, then run tests/torch/ut/trainer/test_trainer_config_build.py, test_trainer_branch.py, and test_base_trainer_unchanged.py; done means all three pass and the legacy LLMTrainer path is unchanged.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python
- Domain
- machine-learning
- Issue type
- Feature
- Difficulty
- 4/5
- Estimated time
- 3-5 days
- Activity status
- Active
- Clarity
- Mostly clear
- Newbie friendliness
- 45/100