mindspore-ai / mindspore-ai/hyper-parallel

【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 3 训练组件Configurable化

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

Part 3 — hyper_parallel/components/ 训练组件 Configurable 化

把现散落在 BaseTrainer._build_*trainer/base.py:555-700)的"过程式 build"抽成 Configurable 子类。新组件类是旧 _build_*等价封装,旧 trainer 代码 0 改动。


1. 目标

让 optimizer / lr_scheduler / loss / tokenizer / dataloader / checkpoint / profiler / metrics 都变成 Configurable 组件,每个 __init__(self, config) 接受自己的 Configconfig.build(...) 即构造。

2. 任务边界(新增文件)

新增包目录 hyper_parallel/components/

新文件 旧逻辑出处 内容 对应 torchtitan
optimizer.py base.py:555-604 _build_optimizer OptimizersContainer(Configurable) + OptimizersInBackwardContainer(Configurable)Config(lr, betas, eps, weight_decay, foreach, fused, decay_keywords, loss_aggregation)__init__(self, cfg, model_parts: list[Module]) 内部跑 decay / no_decay 分组逻辑 torchtitan/components/optimizer.py
lr_scheduler.py base.py:606-637 _build_lr_scheduler LRSchedulersContainer(Configurable):cosine + warmup,复制 _lr_lambda 逻辑;Config(warmup_steps, decay_style, lr_min, lr_max)__init__(self, cfg, optimizers, training_steps) torchtitan/components/lr_scheduler.py
loss.py models/qwen3_5/model.py:361-373F.cross_entropy BaseLoss(Configurable) + CrossEntropyLoss(Configurable);可选 ChunkedCELoss;模型 forward 重构为返回 logits,loss 由 loss.build()(logits, labels) 计算 torchtitan/components/loss.py
tokenizer.py llm_trainer.py:72-104 _build_model_assets BaseTokenizer(Configurable) + HuggingFaceTokenizer(Configurable)Config(path, trust_remote_code, pad_token) torchtitan/components/tokenizer.py
dataloader.py base.py:248-405 _build_dataset / collate / dataloader + llm_trainer.py:202-381 BaseDataLoader(Configurable) + HuggingFaceTextDataLoader / DummyDataLoader / PresetPtDataLoaderbuild(dp_world_size, dp_rank, tokenizer, seq_len) 返回 StatefulDataLoader torchtitan/components/dataloader.py
checkpoint.py base.py:1273-1417 _load_weights / _load_hf_safetensors / _load_hyper_dcp + 现 CheckpointCallback CheckpointManager(Configurable):封装 core/distributed_checkpoint.load + state_dict_adapterConfig(output_dir, save_steps, save_async, save_hf_weights, load_path) torchtitan/components/checkpoint.py
profiler.py ProfilerCallback Profiler(Configurable):torch.profiler / mindspore profiler 跨后端封装;Config(enabled, output_dir, wait_steps, warmup_steps, active_steps) torchtitan/components/profiler.py
metrics.py LoggingCallback / TensorBoardCallback / WandbCallback MetricsProcessor(Configurable):tb / wandb / throughput;Config(report_to, log_steps, tensorboard, wandb) torchtitan/components/metrics.py

3. 核心设计点

3.1 组件 Config 字段名 = 旧 yaml 字段名

例如 OptimizersContainer.Config 的字段必须包含 lr / weight_decay / eps / betas / foreach,与 trainer/config.py:195 OptimizerConfig 字段一致 —— 这样 M2 的 _legacy_to_trainer_config 能一对一映射。

3.2 新旧路径共享同一份"算法实现"

M3 的实现就是把 _build_* 函数体搬到组件 __init__;M5 中旧 BaseTrainer._build_optimizer 改为:

def _build_optimizer(self):
    if isinstance(self.args, HyperTrainer.Config):     # 新路径
        self.optimizer = self.args.optimizer.build(model_parts=[self.model])
    else:                                              # 旧路径
        # base.py:555-604 原逻辑保持不变
        ...
3.3 示例:OptimizersContainer
class OptimizersContainer(Configurable):
    @dataclass(kw_only=True, slots=True)
    class Config(Configurable.Config):
        lr: float = 1e-4
        betas: tuple[float, float] = (0.9, 0.999)
        eps: float = 1e-8
        weight_decay: float = 0.01
        foreach: bool | None = None
        fused: bool | None = None
        decay_keywords: tuple[str, ...] = ("bias", "layernorm", "norm", "rmsnorm")
        loss_aggregation: str = "token_weighted"

    def __init__(self, cfg: "OptimizersContainer.Config", *, model_parts: list):
        # 直接搬 base.py:566-600 的逻辑:
        decay_params, no_decay_params = [], []
        seen = set()
        for m in model_parts:
            for n, p in m.named_parameters():
                if not p.requires_grad or id(p) in seen:
                    continue
                seen.add(id(p))
                lname = n.lower()
                bucket = no_decay_params if any(k in lname for k in cfg.decay_keywords) else decay_params
                bucket.append(p)
        param_groups = [
            {"params": decay_params, "weight_decay": cfg.weight_decay},
            {"params": no_decay_params, "weight_decay": 0.0},
        ]
        self._optimizer = torch.optim.AdamW(
            param_groups, lr=cfg.lr, betas=cfg.betas, eps=cfg.eps, foreach=cfg.foreach,
        )

    def step(self): self._optimizer.step()
    def zero_grad(self): self._optimizer.zero_grad()
    # ... 实现 .state_dict / .load_state_dict 等接口

4. 与 torchtitan 接口差异说明

# 差异点 原因
1 CheckpointManager 封装 hyper 的 DCP(core/distributed_checkpoint/load),不是 torch DistributedCheckpointing hyper 有自己的 DCP 实现
2 MetricsProcessor 整合多个 callback —— 现状 hyper 有 13 个 callback(base.py:680)散落 M3 不改 callback 体系,只增 MetricsProcessor 作为"组件式入口",内部丢给 BaseTrainer 走老接口;M8 统一
3 tokenizer / dataloader 字段保留 hyper 现用名(text_key / streaming / num_workers / pin_memory 旧 yaml 兼容
4 dtype 字段用 Literal[str] 而非 torch.dtype tyro 解析 + 跨后端
5 DummyDataLoader / PresetPtDataLoader 是 hyper 独有的两个特殊路径 llm_trainer.py:202-341 已有逻辑,要保留

5. 开发步骤(每个组件 0.5 d)

按依赖顺序:optimizer → lr_scheduler → loss → tokenizer → dataloader → checkpoint → profiler → metrics

每个组件统一流程:

  1. 把旧 _build_<x> 函数体逐行复制<X>Container.__init__
  2. 把对 self.args.train.<x> 的访问改为 self.config.<field>
  3. 加 1 个 1-card UT:构造 Config()build(...) 后与旧 _build_<x> 同种子 + 同输入下产物 bit-exact

6. 验证标准

新建 tests/torch/ut/components/,每个组件 1 个对照测试:

测试 断言要点
test_optimizer.py 用相同 model_parts 喂给旧 _build_optimizerOptimizersContainer.Config(...).build(model_parts=[m]),参数分组、lr / wd / betas / foreach 完全相同
test_lr_scheduler.py 同种子 100 步 LR 序列 bit-exact
test_loss.py 同 logits / labels 下 CrossEntropyLoss.Config().build()(logits, labels)F.cross_entropy(...) 数值一致
test_tokenizer.py tokenize 同一文本输出 token id 列表一致
test_dataloader.py 同种子下 DummyDataLoader / HuggingFaceTextDataLoader / PresetPtDataLoader 产 1 步样本与旧路径完全一致
test_checkpoint.py 同 state_dict 走两条路径保存/加载,张量逐元素一致
test_profiler.py enabled=False 时 build 出的 profiler 是 nullcontext
test_metrics.py 同 metrics 走两条路径,tb / wandb 写入 key 集合一致

通过门槛

  • 8 个 UT 全绿、bit-exact
  • BaseTrainer._build_*base.py:555-7000 修改
  • 存量训练 0 受影响。

7. 工期 & 依赖

工期 4 天(8 组件 × 0.5 d)
依赖 M1(Configurable
并行 与 M2 / M4 完全并行
下游 M5 / M6

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

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 the M1 Configurable implementation, then compare the existing builders in trainer/base.py, llm_trainer.py, and models/qwen3_5/model.py with the proposed files under hyper_parallel/components/. Run the eight component tests in tests/torch/ut/components/ as they are added. Done means all eight tests are bit-exact, old BaseTrainer builders remain unchanged, and existing training is unaffected.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
machine-learning, testing-qa
Issue type
Feature
Difficulty
5/5
Estimated time
3-5 days
Activity status
Active
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.