已开启
【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 2 配置系统 #140
changzherui创建于 5月15日
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M2 配置系统”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M2”
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M2 配置系统”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M2”
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M2 配置系统”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M2 配置系统”
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M2 配置系统”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M2 配置系统”
5月25日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 2 配置系统”,原标题为“【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M2 配置系统”
5月25日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 2 配置系统”,原标题为“【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M2 配置系统”
M2 —
hyper_parallel/config/配置系统1. 目标
--module qwen3_5_v2 --config qwen3_5_v2_4b --training.steps=10)。TrainingConfig / ParallelismConfig / ...)。_legacy_to_trainer_config把旧 yaml 字段映射到新结构,老用户脚本不破坏。2. 任务边界(新增文件)
新增包目录
hyper_parallel/config/:__init__.pyConfigManager / TrainingConfig / ParallelismConfig / ActivationCheckpointConfig / CompileConfig / CommConfig / DebugConfig / Function / TORCH_DTYPE_MAP;re-exportConfigurableconfigs.py@dataclass(kw_only=True, slots=True)torchtitan/config/configs.pyfunction.pyFunction(Generic[R], Configurable):把任意 callable 包成 Configurable,Function.Config(fn=callable),build()返回fn本身torchtitan/config/function.pymanager.pyConfigManager:parse_args(argv=None) -> Configurable.Config。① 若argv[0]起始为--module→ 走 tyro 路径;② 否则走旧 yaml 路径并经_legacy_to_trainer_config收口torchtitan/config/manager.pydtype_map.pyTORCH_DTYPE_MAP: dict[str, dtype],跨后端通过platform.dtype_map取值,避免直接import torchtorchtitan/config/dtype_map.pytyro_rules.pylist[str]逗号分隔等 tyro custom rules;首次tyro.cli调用前由manager.py自动 installtorchtitan/config/tyro_rules.pylegacy_adapter.py_legacy_to_trainer_config(legacy: HyperTrainerConfig) -> HyperTrainer.Config,把旧 yaml 字段映射到新结构3. 字段对照表(旧 yaml ↔ 新
*Config)train.max_stepsconfig.py:295training.stepstrain.global_batch_sizeconfig.py:297training.global_batch_sizetrain.micro_batch_sizeconfig.py:298training.local_batch_sizetrain.seedconfig.py:299debug.seedtrain.init_deviceconfig.py:303training.init_devicetrain.comm_backendconfig.py:304comm.backendtrain.accelerator.dp_shardconfig.py:149parallelism.data_parallel_shard_degreetrain.accelerator.dp_replicateconfig.py:148parallelism.data_parallel_replicate_degreetrain.accelerator.tpconfig.py:150parallelism.tensor_parallel_degreetrain.accelerator.cpconfig.py:151parallelism.context_parallel_degreetrain.accelerator.ppconfig.py:152parallelism.pipeline_parallel_degreetrain.accelerator.epconfig.py:153parallelism.expert_parallel_degreetrain.accelerator.etpconfig.py:154parallelism.expert_tensor_parallel_degreetrain.accelerator.reshard_after_forwardconfig.py:156parallelism.reshard_after_forwardtrain.mixed_precision.enabled / param_dtype / reduce_dtypeconfig.py:172-175training.mixed_precision_param / training.mixed_precision_reducetrain.gradient_checkpointing.activation_checkpointconfig.py:183activation_checkpoint.modetrain.optimizer.lr / weight_decay / eps / betas / lr_warmup_ratio / lr_decay_style / max_grad_normconfig.py:195-205optimizer.lr / wd / eps / betas+lr_scheduler.warmup_steps / decay_style / max_grad_normtrain.checkpoint.*config.py:208-214checkpoint.*train.profile.*config.py:242-250profiler.*train.debug.*config.py:271-281debug.*model.weights_path / tokenizer_path / freeze_modules / config_overridesconfig.py:67-93model_spec.model.<field>+tokenizer.path+training.freeze_modulesdata.type / train_path / max_seq_len / text_key / num_workers / shuffleconfig.py:99-126dataloader.dataset / dataloader.dataset_path / training.seq_len / dataloader.text_key / dataloader.num_workers / dataloader.shuffle4. 核心设计:
ConfigManager.parse_args双入口class ConfigManager: def parse_args(self, argv: list[str] | None = None) -> "Configurable.Config": argv = argv if argv is not None else sys.argv[1:] # 1. 新风格:--module X --config Y [--field=value] if argv and argv[0].startswith("--module"): return self._parse_new(argv) # 2. 旧风格:yaml + dot-path(兼容现有 train_lm.py) return self._parse_legacy(argv) def _parse_new(self, argv): try: import tyro except ImportError as exc: raise ImportError( "New CLI requires `tyro`. Install with: pip install tyro" ) from exc install_tyro_rules() # 解析 --module / --config,从 config_registry 取 recipe head, remaining = _extract_module_config(argv) module_name, config_name = head["module"], head["config"] registry = importlib.import_module( f"hyper_parallel.models.{module_name}.config_registry" ) default = getattr(registry, config_name)() # HyperTrainer.Config 实例 from hyper_parallel.trainer.trainer_config import HyperTrainer return tyro.cli(HyperTrainer.Config, args=remaining, default=default) def _parse_legacy(self, argv): from hyper_parallel.trainer.config import parse_args, HyperTrainerConfig from hyper_parallel.config.legacy_adapter import _legacy_to_trainer_config legacy = parse_args(HyperTrainerConfig) # 现有 :547 逻辑 return _legacy_to_trainer_config(legacy)5. 与 torchtitan 接口差异说明
Literal["bfloat16", "float32", "float16"]而非torch.dtypetorch.dtype;运行时TORCH_DTYPE_MAP[s]解析ConfigManager.parse_args同时支持新 / 旧两条入口tyro列为可选依赖(requirements.txt加tyro>=0.9.0; extra == "config-cli")local_rank字段保留 hyper 旧语义(从LOCAL_RANK环境变量取),新版放comm.local_rankModelSpec字段在HyperTrainer.Config中标tyro.conf.SuppressModelSpec持 callable / dataclass,tyro 不能从 CLI 解析6. 开发步骤
configs.py—— 7 个 dataclass,按字段对照表逐项落地,写__post_init__校验(如local_batch_size > 0、tensor_parallel_degree >= 1)。function.py+dtype_map.py。manager.py的双入口 +tyro_rules.py。legacy_adapter.py。7. 验证标准
新建
tests/torch/ut/config/:test_configs.pyexamples/yaml/train_qwen3_5_*.yaml对齐;非法值(如tensor_parallel_degree=-2)__post_init__raisetest_function.pyFunction.Config(fn=lambda x: x+1).build()(3) == 4test_manager_legacy_yaml.pyparse_args(HyperTrainerConfig),②ConfigManager().parse_args([yaml_path]);后者过_legacy_to_trainer_config后字段对齐test_manager_new_cli.pytests/torch/ut/config/_test_dummy/config_registry.py写qwen3_5_v2_tiny() -> HyperTrainer.Config;调ConfigManager().parse_args(["--module", "_test_dummy", "--config", "qwen3_5_v2_tiny", "--training.steps=5"]),断言覆盖生效test_legacy_roundtrip.pylegacy → new → legacy后字段值与原对象assertdataclass_equaltest_tyro_missing.pypip install tyro" 的ImportError通过门槛:
parse_args(HyperTrainerConfig)(trainer/config.py:547)行为 0 改动。hyper_parallel.config可独立 import;hyper_parallel/__init__.py不导出。8. 工期 & 依赖
Configurable)