已开启
RFC: Training Input Design #281
TonghanZhang创建于 7月7日
7月7日 添加了label:feature
7月7日 修改了issue 的描述
7月7日 修改了issue 的描述
7月7日 修改了issue 的描述
7月7日 修改了issue 的描述
7月7日 修改了issue 的描述
7月7日 修改了issue 的描述
7月7日 修改了issue 的描述
7月7日 修改了issue 的描述
7月9日 关联了pull request:Align DModule sharding flow with TorchTitan-style config
7月18日 修改了issue 的描述
7月18日 修改标题为 “[Feature]: Training Input Design”,原标题为“[Feature]: 增加 Training Input:Python preset、有类型 CLI 覆盖与运行配置记录”
7月18日 修改标题为 “[Feature]: Training Input Design”,原标题为“[Feature]: 增加 Training Input:Python preset、有类型 CLI 覆盖与运行配置记录”
7月18日 修改标题为 “RFC: Training Input Design”,原标题为“[Feature]: Training Input Design”
7月18日 修改了issue 的描述
支持 YAML 生成统一的 typed
TrainerConfig1. 基本信息
config、trainer、训练组件2. 背景
当前 Hyper-Parallel 的训练输入路径为:
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 前完成以下检查:TrainerConfig字段类型。解析成功后,
TrainerConfig的一级字段直接挂载组件 Config、参数 Config 或ModelSpec。3. 目标和非目标
3.1 目标
TrainerConfig,替换固定的HyperTrainerConfig(model, data, train)结构。_target_选择 Config 类或 factory。TrainerConfig字段、target 参数签名和返回类型完成校验。TrainerConfig上应用 typed CLI dotted override。TrainerConfig的全部一级字段提供明确的 Config 接口。ModelSpec。3.2 非目标
BaseTrainer._build_*的运行构造职责,也不执行 train step 或 checkpoint load。HyperTrainerConfig、model registry 和 discovery 路径在训练入口完成迁移后统一删除。4. 相关实现参考
Trainer.Config直接挂载 model spec、dataloader、optimizer、scheduler、loss、checkpoint 等 typed 组件;Trainer 按运行依赖调用各组件build()TrainerConfig和组件 Config 边界,同时保留 YAML 入口_target_可以导入并调用对象;数据路径正在转向 typed Dataset、Dataloader 和 Collator 的受控构造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 的对象所有权定义:TrainerConfig字段类型modelModelSpectokenizerTokenizer.Configbuild()创建 tokenizerdataloaderDataLoader.ConfigoptimizerOptimizer.Configbuild()lr_schedulerLRScheduler.Configbuild()lossLoss.Configbuild()loss callabletrainingTrainingConfigparallelismParallelismConfigcheckpointCheckpoint.Configbuild()checkpoint manageractivation_checkpointActivationCheckpointConfigmetricsMetrics.Configbuild()指标处理组件profilerProfiler.Configbuild()profilervalidatorValidator.Configbuild()debugDebugConfigcompileCompileConfigcommCommConfigmodel对应 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
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.configTrainerConfig、TrainingConfig、ParallelismConfig及其他参数 Confighyper_parallel.config.resolverTrainerConfighyper_parallel.config.managerConfig类型和build()接口create_model_spec(...) -> ModelSpectrain_lm.py/train_vl.pyTrainerConfig并选择对应 Trainer6.4 方案取舍
TrainerConfig新增训练参数或组件时,先由使用该字段的训练流程负责人和所属组件负责人确定:
TrainingConfig、ParallelismConfig等参数类,还是新的一级组件;build();新增 optimizer、loss 等实现时,应在所属组件模块增加对应 Config/build,而不是向 YAML 开放未声明字段。
7. 组件依赖
TrainerConfig字段定义ModelSpecfactorymodel字段无法解析BaseTrainer._build_*完整能力需要
TrainerConfig字段、组件 Config 接口、resolver 和 CLI override 同时可用。本阶段最小可交付能力是完整 YAML 到 typedTrainerConfig的解析与错误检查。8. 约束与兼容性
resolve_root(raw)生成TrainerConfig;每个一级组件分组通过_target_选择 Config 类或 factoryModelSpec,不创建训练运行对象HyperTrainerConfig、registry/discovery 和BaseTrainer._build_*9. 验证设计
9.1 用例分层
9.2 解析验收
TrainerConfig;各字段对象类型正确TrainerConfig提供正确默认对象_target_optimizer.lrr等未知参数在调用 target 前失败TrainerConfig字段类型时失败9.3 性能 / 显存验证
本阶段不改变训练运行路径,不设置训练性能或显存目标。验证解析耗时不进入 train step 即可。
10. 实现计划
TrainerConfig、一级 target resolver、config manager、typed CLI override、全部组件 Config 接口、内置模型 factoryHyperTrainerConfig、registry/discovery 路径PR1 的完成标准:一份显式覆盖
TrainerConfig全部一级字段的 YAML 能生成完整 typed 配置;具有默认值的字段可以省略;CLI override 可用;所有 target、参数和类型错误均在进入 Trainer 前报告。