mindspore-ai / mindspore-ai/hyper-parallel
RFC: Training Input Design
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 |
| 相关模块 | config、trainer、训练组件 |
| 相关 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 目标
- 提供
TrainerConfig,替换固定的HyperTrainerConfig(model, data, train)结构。 - YAML 中实际提供的一级组件分组分别通过
_target_选择 Config 类或 factory。 - resolver 根据
TrainerConfig字段、target 参数签名和返回类型完成校验。 - 支持在最终
TrainerConfig上应用 typed CLI dotted override。 - 为
TrainerConfig的全部一级字段提供明确的 Config 接口。 - 内置模型 package 通过 factory 返回现有
ModelSpec。
3.2 非目标
- 本阶段不创建 Trainer、model、optimizer、dataloader 等训练运行对象。
- 本阶段不迁移
BaseTrainer._build_*的运行构造职责,也不执行 train step 或 checkpoint load。 - resolver 只解析 YAML 的一级组件分组,在框架/用户输入层面禁止嵌套_target_和未注册的任意新组件的解析。
- Dataset、Collator 和 DataLoader 的运行时依赖与构造顺序由后续数据组件迁移定义。
- Hugging Face 模型 target、resolved-config 记录和任意 callback 列表分别评审,不纳入本阶段接口。
- 旧
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) 返回 TrainerConfig;train_lm 和 train_vl 分别把解析结果交给 LLMTrainer 和 VLTrainer。
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 |
定义 TrainerConfig、TrainingConfig、ParallelismConfig 及其他参数 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 接口并确定所有权 | 是 | 配置字段、组件接口和错误检查边界明确 |
新增训练参数或组件时,先由使用该字段的训练流程负责人和所属组件负责人确定:
- 字段属于已有
TrainingConfig、ParallelismConfig等参数类,还是新的一级组件; - 运行对象由哪个组件
build(); - 是否需要修改 Trainer 的构造顺序;
- 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
- 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.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