mindspore-ai / mindspore-ai/hyper-parallel

【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 5 Trainer 接缝层

Open
#295 0 comments 0 reactions 0 assignees View on GitHub

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 方法 重构后
_setupbase.py:124-192 抽出 _helper_setup_distributed(config) 纯函数;原方法变为薄包装
_post_parallelizebase.py:425-462 抽出 _helper_post_parallelize(model, init_device, weights_path, mp_cfg, ...)
_materialize_and_init_shardsbase.py:1183-1222 已经是 helper 风格,保持不变
_build_optimizerbase.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_specAnnotated[..., 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.pyHyperTrainer.Config,引用 M2 / M3 的所有 Config。
  3. Step 3(1 d):写 hyper_trainer.pyHyperTrainer.__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

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

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. 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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.