mindspore-ai / mindspore-ai/hyper-parallel

RFC: Training Input Design

Open
#668 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

支持 YAML 生成统一的 typed TrainerConfig

1. 基本信息

项目 内容
作者 Tonghan Zhang
相关模块 configtrainer、训练组件
相关 issue / PR 本 Issue;实现 PR 待创建
适用后端 本特性的配置解析层与后端无关

2. 背景

当前 Hyper-Parallel 的训练输入路径为:

YAML / CLI
→ parse_args(HyperTrainerConfig)
→ HyperTrainerConfig(model, data, train)
→ discover_model_spec(model.name)
→ LLMTrainer / VLTrainer
→ BaseTrainer._build_*

HyperTrainerConfig 是固定的三层参数树。当前 YAML 中没有 _target_parse_args() 只按 model / data / train 的 dataclass 字段读取参数。具体模型通过 model.name 和 registry/discovery 选择,dataloader、optimizer、checkpoint 等运行对象由 Trainer 和 BaseTrainer._build_* 选择并创建。

本 RFC 新增一级组件 _target_,让 YAML 可以选择具体 Config 类或 factory。引入该接口后,resolver 需要在进入 Trainer 前完成以下检查:

  • target 路径存在且目标可调用;
  • target 参数名合法且必填参数齐全;
  • 参数值符合 target 签名;
  • target 返回值符合对应的 TrainerConfig 字段类型。

解析成功后,TrainerConfig 的一级字段直接挂载组件 Config、参数 Config 或 ModelSpec

本 RFC 要解决的问题:YAML 在进入 Trainer 前生成完整、可校验的 typed TrainerConfig。而非以参数树的形式打包进入LLMTrainer,在主流程里面选择组件的构造顺序与构造逻辑。
完成后的成功标准:完整配置可以解析为 TrainerConfig,所有输入错误在 Trainer 构建前报告。

3. 目标和非目标

3.1 目标
  1. 提供 TrainerConfig,替换固定的 HyperTrainerConfig(model, data, train) 结构。
  2. YAML 中实际提供的一级组件分组分别通过 _target_ 选择 Config 类或 factory。
  3. resolver 根据 TrainerConfig 字段、target 参数签名和返回类型完成校验。
  4. 支持在最终 TrainerConfig 上应用 typed CLI dotted override。
  5. TrainerConfig 的全部一级字段提供明确的 Config 接口。
  6. 内置模型 package 通过 factory 返回现有 ModelSpec
3.2 非目标
  1. 本阶段不创建 Trainer、model、optimizer、dataloader 等训练运行对象。
  2. 本阶段不迁移 BaseTrainer._build_* 的运行构造职责,也不执行 train step 或 checkpoint load。
  3. resolver 只解析 YAML 的一级组件分组,在框架/用户输入层面禁止嵌套_target_和未注册的任意新组件的解析。
  4. Dataset、Collator 和 DataLoader 的运行时依赖与构造顺序由后续数据组件迁移定义。
  5. Hugging Face 模型 target、resolved-config 记录和任意 callback 列表分别评审,不纳入本阶段接口。
  6. HyperTrainerConfig、model registry 和 discovery 路径在训练入口完成迁移后统一删除。

4. 相关实现参考

来源 做法 限制 对本 RFC 的影响
TorchTitan Trainer.Config 直接挂载 model spec、dataloader、optimizer、scheduler、loss、checkpoint 等 typed 组件;Trainer 按运行依赖调用各组件 build() 用户入口主要是 Python 配置函数,不是 Hyper 现有的 YAML Hyper 采用统一根 TrainerConfig 和组件 Config 边界,同时保留 YAML 入口
AutoModel YAML _target_ 可以导入并调用对象;数据路径正在转向 typed Dataset、Dataloader 和 Collator 的受控构造 1. 嵌套Target和任意字段,参数的挂载:框架引入NodeConfig做类型暂存/n 2. 用户使用的时候参数/组件-训练拉起之间界限不明确,存在耦合的组件可以在yaml层面嵌套/在组件文件中耦合...框架不易维护 Hyper 只解析一级组件 target;组件内部组合由所属 Config/build 管理

5. 对外接口

5.1 根配置

Python 代码直接使用导入后的 TrainerConfig

from hyper_parallel.trainer.config import TrainerConfig


def resolve_root(raw: dict) -> TrainerConfig:
    ...

resolve_root(raw) 返回 TrainerConfigtrain_lmtrain_vl 分别把解析结果交给 LLMTrainerVLTrainer

TrainerConfig 的一级字段参考 TorchTitan 的组件边界,并按 Hyper 的对象所有权定义:

YAML 字段 TrainerConfig 字段类型 接口职责
model ModelSpec model factory 组织模型构造、并行和 checkpoint adapter 描述
tokenizer Tokenizer.Config 保存 tokenizer 参数;运行阶段通过 build() 创建 tokenizer
dataloader DataLoader.Config 保存数据管线配置;运行阶段负责 Dataset、Collator、Sampler 和 DataLoader 的受控构造
optimizer Optimizer.Config 保存 optimizer 参数;运行阶段接收模型参数并 build()
lr_scheduler LRScheduler.Config 保存 scheduler 参数;运行阶段接收 optimizer 等依赖并 build()
loss Loss.Config 保存 loss 参数;运行阶段 build() loss callable
training TrainingConfig 保存训练循环参数,不创建运行对象
parallelism ParallelismConfig 保存并行拓扑和策略参数,不创建运行对象
checkpoint Checkpoint.Config 保存 save/load 参数;运行阶段 build() checkpoint manager
activation_checkpoint ActivationCheckpointConfig 保存激活重计算策略参数
metrics Metrics.Config 保存指标和日志参数;运行阶段 build() 指标处理组件
profiler Profiler.Config 保存 profiler 参数;运行阶段 build() profiler
validator Validator.Config 保存验证参数;运行阶段接收 dataloader、loss 等依赖并 build()
debug DebugConfig 保存确定性、数值检查等调试参数
compile CompileConfig 保存编译策略参数
comm CommConfig 保存通信初始化和超时参数

model 对应 TorchTitan 根配置中的 model_spec;Hyper 的 YAML 保留 model 作为用户字段,解析结果类型为 ModelSpec

具有默认值的字段可以在普通训练 YAML 中省略。YAML 中一旦显式提供某个一级字段,该分组就必须包含 _target_

5.2 YAML 示例
model:
  _target_: hyper_parallel.models.qwen3_5.create_model_spec
  weights_path: /path/to/Qwen3.5-0.8B-Base

training:
  _target_: hyper_parallel.trainer.config.TrainingConfig
  max_steps: 100
  global_batch_size: 8

parallelism:
  _target_: hyper_parallel.trainer.config.ParallelismConfig
  tp: 2

optimizer:
  _target_: hyper_parallel.components.optimizer.AdamW.Config
  lr: 0.0002
  weight_decay: 0.1

loss:
  _target_: hyper_parallel.components.loss.CausalLMLoss.Config
  ignore_index: -100

这是最小接口示例。完整解析测试需要显式覆盖 5.1 节列出的全部一级字段。

5.3 CLI dotted override
python -m hyper_parallel.train_lm \
  --config-file configs/qwen3_5.yaml \
  --training.max_steps 200 \
  --parallelism.tp 4 \
  --optimizer.lr 0.0001

resolver 先生成 typed TrainerConfig,再应用 CLI override。CLI 字段必须存在于最终配置类型中,值必须符合字段类型。

5.4 模型 factory

内置模型 package 提供公开 factory:

def create_model_spec(...) -> ModelSpec:
    return ModelSpec(
        name="qwen3_5",
        build_model_fn=_build,
        parallelize_fn=parallelize_qwen3_5,
        state_dict_adapter=Qwen3_5StateDictAdapter,
    )

model._target_ 指向该 factory。resolver 只得到 ModelSpec,不创建模型或加载权重。

6. 方案设计

6.1 总体流程

解析阶段只创建 Config、参数类和 ModelSpec。运行对象由后续 Trainer 与组件 build() 路径创建。

6.2 关键逻辑
def resolve_component(node, path):
    target = import_target(require(node, "_target_", path))
    args = {name: value for name, value in node.items() if name != "_target_"}

    check_fields(target, args, path)
    check_types(target, args, path)

    result = target(**args)
    check_factory_result(target, result, path)
    return result


def resolve_root(raw) -> TrainerConfig:
    check_fields(TrainerConfig, raw, path="$")
    components = {
        name: resolve_component(node, path=name)
        for name, node in raw.items()
    }
    check_types(TrainerConfig, components, path="$")
    return TrainerConfig(**components)

resolver 不维护组件名称列表;一级字段集合由 TrainerConfig 声明。target 内部参数直接传给当前 target,不继续扫描嵌套 _target_

6.3 代码改动点
模块 改动内容 是否影响已有行为
hyper_parallel.trainer.config 定义 TrainerConfigTrainingConfigParallelismConfig 及其他参数 Config 新增 typed 配置接口;运行入口切换前不改变训练行为
hyper_parallel.config.resolver 导入一级 target,校验参数和返回类型,构造 TrainerConfig 新增解析路径
hyper_parallel.config.manager 读取 YAML,调用 resolver,应用 typed CLI override 新增统一配置入口
optimizer、scheduler、loss、dataloader、checkpoint 等组件模块 提供稳定的 Config 类型和 build() 接口 本阶段只定义接口,不迁移运行构造
内置模型 package 提供 create_model_spec(...) -> ModelSpec 保留现有模型构造、并行和 adapter 所有权
train_lm.py / train_vl.py 后续运行迁移时接收 TrainerConfig 并选择对应 Trainer 本阶段保留现有运行路径
配置解析测试 覆盖完整字段、错误路径和 CLI override 不启动 Trainer
6.4 方案取舍
方案 优点 缺点 是否选择 原因
YAML → 通用动态节点 → typed Config 能表达任意嵌套 target 同时维护动态节点和 typed Config;运行依赖边界不清晰 Hyper 不需要第二套中间配置树
YAML 一级 target → TrainerConfig 入口保持 YAML;最终结果只有一棵 typed 配置树 新组件必须先定义 Config 接口并确定所有权 配置字段、组件接口和错误检查边界明确

新增训练参数或组件时,先由使用该字段的训练流程负责人和所属组件负责人确定:

  1. 字段属于已有 TrainingConfigParallelismConfig 等参数类,还是新的一级组件;
  2. 运行对象由哪个组件 build()
  3. 是否需要修改 Trainer 的构造顺序;
  4. Config 字段、YAML 示例和验证用例如何同步更新。

新增 optimizer、loss 等实现时,应在所属组件模块增加对应 Config/build,而不是向 YAML 开放未声明字段。

7. 组件依赖

依赖组件 强依赖 / 弱依赖 当前状态 未 ready 时本阶段能力
TrainerConfig 字段定义 强依赖 本 RFC 定义 无法校验 YAML 一级字段和组件返回类型
各组件 Config 接口 强依赖 本阶段补齐接口 对应字段不能进入完整解析测试
内置模型 ModelSpec factory 强依赖 复用现有 ModelSpec 内容并增加公开 factory model 字段无法解析
BaseTrainer._build_* 弱依赖 现有运行路径可用 不影响配置解析;运行迁移后续完成
PT / MS 后端 不涉及 解析层共用 Python 类型 本阶段不创建后端运行对象

完整能力需要 TrainerConfig 字段、组件 Config 接口、resolver 和 CLI override 同时可用。本阶段最小可交付能力是完整 YAML 到 typed TrainerConfig 的解析与错误检查。

8. 约束与兼容性

类型 内容
配置边界 resolve_root(raw) 生成 TrainerConfig;每个一级组件分组通过 _target_ 选择 Config 类或 factory
target 范围 target 只创建 Config、参数类或 ModelSpec,不创建训练运行对象
嵌套结构 resolver 不递归实例化 target 参数;组件内部组合由所属 Config/build 管理
已有训练行为 运行入口迁移前继续使用旧 HyperTrainerConfig、registry/discovery 和 BaseTrainer._build_*
迁移方式 选定一个内置模型完成组件运行迁移后,统一删除旧配置和 discovery 路径,不保留兼容别名
性能与显存 解析只发生在启动阶段,不改变训练计算、通信、吞吐或显存路径
PT / MS 差异 解析行为一致;后端差异由具体组件 Config/build 所有者处理

9. 验证设计

9.1 用例分层
用例级别 数量 覆盖内容 通过标准
UT 参数化覆盖 target 导入、字段名、必填参数、参数类型、返回类型、CLI override 所有成功与失败路径符合字段路径和类型约束
9.2 解析验收
用例 通过标准
完整 YAML 显式覆盖 5.1 节全部一级字段并生成 TrainerConfig;各字段对象类型正确
默认字段 省略具有默认值的一级字段时,由 TrainerConfig 提供正确默认对象
未知一级字段 在调用任何 target 前失败,错误包含根字段路径
缺少 _target_ 显式提供的一级组件分组立即失败,错误包含字段路径
target 导入 路径不存在或目标不可调用时失败,错误包含字段路径
参数名称 optimizer.lrr 等未知参数在调用 target 前失败
必填参数 缺少 target 必需参数时指出参数名和字段路径
参数类型 参数值不符合 target 参数注解时失败
factory 返回类型 target 返回值不符合 TrainerConfig 字段类型时失败
typed CLI 可覆盖具体组件字段;未知字段和值类型错误立即失败
解析副作用 解析过程不创建 Trainer、model、optimizer、dataloader 或 checkpoint manager
9.3 性能 / 显存验证

本阶段不改变训练运行路径,不设置训练性能或显存目标。验证解析耗时不进入 train step 即可。

10. 实现计划

PR 内容 依赖 验证
PR1 TrainerConfig、一级 target resolver、config manager、typed CLI override、全部组件 Config 接口、内置模型 factory 配置解析 UT
PR2 选择 Llama 3 或 Qwen 3.5,迁移组件构建和训练入口,删除旧 HyperTrainerConfig、registry/discovery 路径 PR1 UT + Level0
后续 PR 按模型或训练流程补充新组件,并由流程负责人和组件负责人共同确定字段所有权与构造顺序 PR2 对应组件 UT + 训练验证

PR1 的完成标准:一份显式覆盖 TrainerConfig 全部一级字段的 YAML 能生成完整 typed 配置;具有默认值的字段可以省略;CLI override 可用;所有 target、参数和类型错误均在进入 Trainer 前报告。

schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 281
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/281

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.config, hyper_parallel.config.resolver, and hyper_parallel.config.manager, then review the component Config interfaces and existing configuration parsing tests. Trace how YAML and typed CLI overrides currently enter train_lm.py and train_vl.py. Done means a complete YAML input resolves to a typed TrainerConfig, validation errors occur before Trainer construction, and parsing tests cover defaults, invalid targets, types, return values, and overrides.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
tooling
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Active
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.