已开启
【RFC】hyper-parallel 支持torchtitan 风格 Module/Config (总览) #138
changzherui创建于 5月15日
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config (总览)”,原标题为“hyper-parallel 引入 torchtitan 风格 Module/Config 体系(总览)”
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config (总览)”,原标题为“hyper-parallel 引入 torchtitan 风格 Module/Config 体系(总览)”
5月15日 修改了issue 的描述
5月15日 修改了issue 的描述
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config (总览)”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config (总览)”
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config (总览)”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config (总览)”
5月16日 修改了issue 的描述
5月16日 修改了issue 的描述
5月25日 修改了issue 的描述
5月25日 修改了issue 的描述
5月25日 修改了issue 的描述
5月25日 修改了issue 的描述
5月25日 修改了issue 的描述
5月25日 修改了issue 的描述
1. 整体任务
把现有"裸
nn.Module+ 字符串_tp_plan+ 过程式parallelize_<name>()"路径,扩展为 torchtitan 风格的:3 条硬约束:
LLMTrainer(args)路径必须 100% 保留至少一个版本周期。import torch.nn/import torch,必须经platform = get_platform();Module基类是(platform.Module, Configurable)。Configurable / Module / ShardingConfig / NamedPlacement / ModelSpec / ConfigManager)完全照抄;不一致的地方在每个模块文档显式列出。flowchart TB subgraph entry [入口层] run_train["run_train.sh / torchrun"] train_py["train.py"] CM["ConfigManager"] end subgraph config_layer [配置层] CR["config_registry.py\nllama3_8b() → Trainer.Config"] TC["Trainer.Config\n(training/parallelism/checkpoint/...)"] MS["ModelSpec\n(model + parallelize_fn + pipelining_fn)"] MC["BaseModel.Config\n(嵌套 Module.Config 树)"] end subgraph runtime [运行时] Trainer["Trainer"] PD["ParallelDims\n(DeviceMesh)"] MP["model.parallelize(parallel_dims)"] FSDP["fully_shard / apply_fsdp"] end run_train --> train_py --> CM CM --> CR --> TC TC --> MS --> MC CM -->|"config.build()"| Trainer Trainer --> PD Trainer -->|"parallelize_fn"| MP --> FSDPconfig驱动module的完整生命周期
① 定义 Config 树(model_registry / llama3_configs)
↓
② update_from_config() — 在 Config 树上填充/修改
(seq_len → rope.max_seq_len,set_llama3_sharding_config() → 各层 sharding_config)
↓
③ meta device 上 root Config.build() — 递归建 Module 树(无真实内存)
↓
④ verify_module_protocol() — 检查 Module 树每个节点都是 Module 子类
↓
⑤ parallelize_fn → model.parallelize(parallel_dims)
— 按 Module 树上的 _sharding_config 做 DTensor 分片
↓
⑥ to_empty(GPU) + init_states() — 按 Module 树上的 _param_init 初始化权重
2. 模块划分
按"代码依赖 + 可独立交付"切成 9 个相互独立的模块:
Configurable / Module / ShardingConfig / MeshAxisName / BaseModel / ModelSpec v2ConfigManager+ tyro CLI + 7 个通用*Config+ 新旧 yaml 适配Linear / RoPE / GQAttention / FeedForward / MoE / Decoder+decoder_sharding声明式助手HyperTrainer.Config / HyperTrainer(BaseTrainer),BaseTrainer.__init__加新分支config_registry.py+ 改scripts/train_lm.py(约 5 行)qwen3_5_moe_v2 / qwen3_vl_moe_v2;切换默认 spec;清理 legacycore/dtensor/local_map.py包装;按需触发3. 模块依赖关系图
关键路径(必须串行):P1 → P4 → P5 → P6 → P7,最快 ~17 天。
并行机会:P1 完成后,P2 / P3 / P4 / P9 可同时推进。4 名工程师并行最快 ~12 天到 P7 完成。
4. 各模块一句话职责
Part 1 —
hyper_parallel/protocols/(协议基石)定义
Configurable / Configurable.Config(__init_subclass__自动绑 owner,Config.build()即构造),以及Module(platform.Module, Configurable)三件套(init_states / parallelize(mesh) / from_nn_module)、ShardingConfig / NamedPlacement / MeshAxisName、BaseModel / ModelConfigConverter、新版ModelSpec。纯新增,零依赖。Part 2 —
hyper_parallel/config/(配置系统)基于 P1 的
Configurable,提供:ConfigManager.parse_args()双入口(--module走 tyro / 否则走旧 yaml 适配)、TrainingConfig / ParallelismConfig / ActivationCheckpointConfig / CompileConfig / CommConfig / DebugConfig等通用 dataclass、Function(Configurable)、TORCH_DTYPE_MAP、_legacy_to_trainer_config旧字段映射。tyro列为可选依赖。Part 3 —
hyper_parallel/components/(训练组件 Configurable 化)把现散落在
BaseTrainer._build_*(trainer/base.py:555-700)的"过程式 build"抽成Configurable子类:OptimizersContainer / LRSchedulersContainer / BaseLoss + CrossEntropyLoss / BaseTokenizer + HuggingFaceTokenizer / BaseDataLoader + HuggingFaceTextDataLoader + DummyDataLoader / CheckpointManager / Profiler / MetricsProcessor。算法实现直接搬运旧_build_*函数体,保证 bit-exact 等价。Part 4 —
hyper_parallel/models/common/(通用模型组件库)把
hyper_parallel/models/modules/现有裸nn.Module组件重写为 Module 协议版:Linear / Embedding / RMSNorm / Qwen3_5RMSNorm / RoPE / GQAttention / FeedForward / MoE / TransformerBlock / Decoder,加param_init.py、decoder_sharding.py(9 个声明式 sharding 助手)。旧models/modules/保留供老路径继续用。Part 5 —
trainer/接缝层新增
trainer_config.py(HyperTrainer.Config顶层 Configurable 树)+hyper_trainer.py(HyperTrainer(BaseTrainer)新子类,11 步__init__全走cfg.build())。BaseTrainer.__init__只新增一个分支,把现有 13 步_build_*抽 helper,旧LLMTrainer走旧 helper、新HyperTrainer走 helper 加 build —— 算法零分叉。Part 6 —
models/qwen3_5_v2/(首个模型迁移)按 P1 / P4 重写
Qwen3_5Model(Decoder)+Qwen3_5TransformerBlock,set_qwen3_5_sharding_config(),parallelize_qwen3_5_v2()(顺序 CP →model.parallelize(tp_mesh)→ AC → FSDP),Qwen3_5_v2StateDictAdapter(BaseStateDictAdapter),register_spec("qwen3_5_v2", ...)。旧models/qwen3_5/完全保留。Part 7 — CLI 入口打通
写
models/qwen3_5_v2/config_registry.py(qwen3_5_v2_debugmodel() / _4b() / _4b_tp2_fsdp4()等 recipe,每个返回HyperTrainer.Config);改scripts/train_lm.py入口(约 5 行):mgr = ConfigManager() config = mgr.parse_args() # 自动识别 --module / yaml trainer = config.build() # Configurable.Config.build() → HyperTrainer(config) trainer.train()Part 8 — 其他模型迁移 + 切换默认
按 Part 6 模式迁移
qwen3_5_moe_v2 / qwen3_vl_moe_v2;待全部 v2 稳定 ≥ 1 周后切换默认:register_spec("qwen3_5", ...)切到 v2 实现,旧实现挪到_legacy/加@deprecated;hyper_parallel/__init__.py暴露新 API;写docs/zh/module_protocol.md。Part 9 —
local_map扩展(可选)新增
hyper_parallel/core/dtensor/local_map.py(torch 后端转发torch.distributed.tensor.experimental.local_map,mindspore 后端 raise);改Module.parallelize在sharding_config.local_map is not None时启用;新增set_gqa_inner_attention_local_map(...)助手。只有当某个模型必须走q/k/vhead-shard 路径时才触发。5. 风险与红线
hyper_parallel/__init__.py在 Part 8.3 之前不暴露任何新 API,避免阶段未稳定就成为公共面。LLMTrainer(args)路径在 Part 8.3 切换前 0 修改,保护现网用户。tyro是新增依赖,列入可选 extras;Part 2 必须保证缺失时报清晰错误。platform;Part 9 在 mindspore 后端用 capability flag 控制。6. 必要的不一致(与 torchtitan 接口的差异汇总)
Module继承(platform.Module, Configurable)而非(nn.Module, Configurable)distribute_tensor用hyper_parallel.core.dtensor.distribute_tensor(core/dtensor/dtensor.py:409)MeshAxisName字面量优先 hyper 现用名(FSDP="fsdp"),同时收DP_SHARD="dp_shard"trainer/parallel_dims.py已建的 mesh 命名对齐local_map第一阶段未实现(M9 才落)ConfigManager同时支持新 tyro / 旧 yaml 入口Literal["bfloat16", "float32", "float16"]torch.dtype,且要跨后端ModelSpecv2 兼容 v1:build_model_fn / model: BaseModel.Config二选一ModelSpec字段标tyro.conf.SuppressMetricsProcessor做组件式入口,但不重写 callback;M8 再统一7. issue任务清单
00_overview.mdPart 1_protocolsPart 2_configPart 3_componentsPart 4_models_commonPart 5_trainer_gluePart 6_qwen3_5_v2Part 7_cli_entrypointPart 8_migrate_othersPart 9_local_map