mindspore-ai / mindspore-ai/hyper-parallel

[RFC] end2end trainer w/ Qwen3-30B training demo

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

【RFC】HyperModels Trainer 统一编排与端到端训练闭环

0. 基本信息

项目 内容
特性名称 HyperModels Trainer 统一编排与端到端训练闭环
开发分支 trainer-dev
Trainer 范围 hyper_models/trainer/**
端到端示例 examples/training_demo/**
主要后端 PyTorch
当前验收拓扑 8 卡 Ascend NPU,HCCL,FSDP dp_shard_size=8
当前示例模型 Qwen3-30B-A3B,随机初始化权重
当前示例数据 WikiText-2 train split,固定长度 1024 tokens

1. 背景

HyperParallel 已经具备 DeviceMesh、DTensor、声明式分片、FSDP/HSDP 以及多种组合并行能力,但完整训练
还需要解决另一组编排问题:

  1. 如何把 YAML 配置解析为有类型、可校验的运行时组件;
  2. 如何只创建一次分布式拓扑,并让模型、数据、loss 和指标共享同一语义;
  3. 如何保证模型完成构建、并行化和参数物化后,再创建 optimizer 等依赖参数身份的组件;
  4. 如何把 global batch 拆成每个 DP rank 的 optimizer batch 和多个 micro batch;
  5. 如何在梯度累积下正确控制 FSDP reshard、gradient sync、梯度裁剪和 optimizer step;
  6. 如何把日志、进度、环境指标、评估触发和内存清理从核心训练步中拆出;
  7. 如何提供一个真实模型、真实数据和多卡 FSDP 的端到端入口,证明上述组件能够串成闭环。

当前分支已经形成以 TrainerConfig + Target + BaseTrainer + TextTrainer + Callback 为核心的实现,并新增
examples/training_demo 作为 Qwen3-30B-A3B/WikiText/FSDP 示例。需要通过 RFC 固化当前职责边界,明确
已实现能力与占位接口,避免把配置面的可表达性误认为运行时支持。


2. 目标与非目标

2.1 目标

本 RFC 固化以下能力:

  1. 使用一个 YAML 文件描述模型、tokenizer、数据变换、dataset、collator、dataloader、loss、optimizer、
    scheduler、并行拓扑和 callback cadence。
  2. 通过 parse_training_args() 解析配置,并支持 --section.field=value 形式的强类型 CLI 覆盖。
  3. 通过 Target 延迟构建需要运行时依赖的组件,配置解析阶段不提前创建 model、dataset 或 optimizer。
  4. TextTrainer 作为当前文本训练规范入口,按显式依赖顺序组装 BaseTrainer 的各阶段。
  5. 分布式初始化和 DistributedSetup/MeshContext 只创建一次,并注入 AutoModel 构建路径。
  6. 模型由 HyperAutoModelForCausalLM.from_pretrained() 原子完成构造、分片、FSDP 包裹和物化;Trainer
    不再执行第二次并行化。
  7. DataLoader 每次迭代返回 list[dict[str, Any]],该列表表示一个 optimizer step 内的全部 micro batches。
  8. 训练步完成 token 统计、forward/backward、FSDP 累积同步控制、梯度裁剪、optimizer step、scheduler
    step 和指标发布。
  9. 使用 callback 承担环境指标、结构化日志、tqdm、评估触发占位和周期性内存清理。
  10. 使用唯一的 examples/training_demo 端到端用例验收 Trainer 完整链路。
2.2 非目标

以下能力不属于本期验收范围:

  • 完整 validation dataloader 和评估循环;
  • checkpoint 保存、恢复、断点续训和配置持久化;
  • PP 调度和 pipeline stage 间通信;
  • TP、CP、EP 及其与 FSDP 的组合并行验收;
  • 生产级 PEFT、QAT、FP8、参数冻结和真实预训练权重加载;
  • packed sequence、动态 batching、多数据源、VLM 或 RL 数据契约;
  • loss 数值对拍、性能基线、显存收益或收敛性结论;
  • examples/training_demo 之外的验证用例。

3. 当前代码结构与职责

hyper_models/trainer/
├── config.py
│   ├── TrainerConfig 及各配置 dataclass
│   ├── Target:延迟构建协议
│   └── save_configs:当前为空实现
├── base.py
│   ├── BackgroundPrefetcher / HyperIter
│   └── BaseTrainer:组件构建、训练步、训练循环、资源销毁
├── text_trainer.py
│   └── TextTrainer:文本训练规范入口,以 composition 方式复用 BaseTrainer
└── callbacks/
    ├── base.py:TrainerState 与 Callback 生命周期
    ├── environ_meter_callback.py:指标生产与跨 rank 归约
    ├── logging_callback.py:rank 0 结构化日志
    ├── tqdm_callback.py:rank 0 进度条
    ├── evaluate_callback.py:评估触发占位
    ├── garbage_collection_callback.py:GC 与 device cache 清理
    └── temp_log_callback.py:TqdmCallback 的兼容别名

Trainer 目录只负责训练编排。具体能力的唯一所有者如下:

能力 所有者 Trainer 责任
YAML 读取、target 导入、类型校验、CLI override hyper_models/config/ 消费解析后的 TrainerConfig
DeviceMesh 与进程组 components/distributed/infrastructure.py 初始化一次并保存引用
模型构建和并行化 _transformers/components/distributed/ 调用 config.model.build(distributed_setup=...)
数据实现 components/data/ 或用户 target 按依赖顺序调用 target
loss 计算与全局归一化 components/loss/ 传入 model output、labels 和 mesh
optimizer/scheduler 实现 components/optim/ 或用户 target 在模型参数身份稳定后构建并逐步调用
训练生命周期 hyper_models/trainer/ epoch、step、micro-step、callback 和销毁

4. 当前支持矩阵

能力 当前状态 说明
强类型 YAML 根配置 支持 未知字段、缺少必填字段和类型不匹配会在解析阶段报错
_target_ 延迟构建 支持 配置参数与运行时参数合并,运行时参数优先
typed CLI override 支持 使用 --field=value,不支持通过 override 更换 _target_
分布式初始化 支持 NPU 上将 NCCL 请求映射到 HCCL;CPU 可回退 Gloo
training.init_device 配置占位 当前 device 由运行环境和 local rank 决定,该字段未参与选择
DP/FSDP mesh 示例路径支持 demo 使用 dp_shard_size=8
TP/CP/EP/PP 配置 配置可表达 不属于本 RFC 的端到端验收范围;PP 当前仍是 stub
AutoModel 原子建模 部分支持 构建、分片、FSDP、物化路径可达;meta 路径真实权重加载未实现
通用 FSDP 模型识别 部分支持 当前 FSDP2Manager 只识别 GPT-2、Llama、Qwen3-MoE 示例结构
tokenizer/data transform/dataset 支持 target 构建 demo 使用 AutoTokenizer、WikiText 和 PlainTextDataTransform
micro-batch contract 支持 一个 dataloader item 固定为非空 list[dict]
background prefetch 支持 单后台线程,保留最近已消费 dataloader state
model-output loss 支持 默认读取 ModelOutput.loss
自定义 loss target 支持 target 必须构建为 torch.nn.Module
token-weighted loss 支持 当前实际路径固定调用 mean_global_loss()
rank-average loss 配置占位 loss_aggregation 可配置,但当前训练步未分派到该语义
梯度累积下 FSDP sync 支持 非最终 micro batch 延迟 gradient sync/all-reduce
梯度裁剪与参数更新 支持 HSDP stream 完成后裁剪,再执行 optimizer/scheduler
训练与环境指标 支持 loss、grad norm、LR、step time、tokens/s、samples 和 device memory
tqdm 与结构化日志 支持 仅 global rank 0 输出
周期性 GC/cache 清理 支持 gc_stepsempty_cache_steps 独立控制
训练中 evaluation 占位 当前只发 warning,不执行验证前向
checkpoint/resume 未闭环 配置存在,save_configs() 为空,Trainer 未注册 checkpointer callback
WandB 配置占位 WandbConfig 存在,但 callback 未接入
plan overhead 配置占位 plan_overhead target 当前未被 Trainer 或 AutoModel 消费
mixed precision/recompute/debug 配置占位或下游预留 Trainer 训练上下文仍为 nullcontext(),NaN/Inf 检查未接入
packed sequence/magi 配置占位 当前文本 demo 不消费
PEFT 部分接线 Trainer 将配置传给 AutoModel,但下游注入仍为 stub

5. 对外接口

5.1 Python 入口

当前规范入口为:

from hyper_models.config.manager import parse_training_args
from hyper_models.trainer.text_trainer import TextTrainer

config = parse_training_args()
trainer = TextTrainer(config)
trainer.train()

BaseTrainer 是实现载体,不作为当前示例的直接用户入口。TextTrainer 通过 composition 持有
BaseTrainer,并显式选择文本训练需要的构建阶段与生命周期。

5.2 CLI 入口
python -m examples.training_demo.train_text \
    examples/training_demo/train.yaml \
    --training.max_steps=10 \
    --optimizer.lr=0.0002

CLI override 规则:

  • 必须使用 --field=value
  • 支持 dataclass 字段和已选 Target 参数的 dotted path;
  • 值先经 YAML scalar 解析,再按目标类型校验;
  • 未知字段、未选择的可选组件和无效类型必须明确报错;
  • 不允许通过 CLI override 修改 _target_
5.3 Target 构建协议

YAML 中包含 _target_ 的节点会被解析为:

Target(
    resolved_callable,
    target_path="package.module.callable",
    **configured_kwargs,
)

运行时调用:

component = target.build(**runtime_kwargs)

最终参数为:

configured kwargs + applicable runtime kwargs

同名运行时参数覆盖配置参数。若 target 不接受 **kwargsTarget.build() 会过滤 target 签名中不存在的
运行时参数,使不同组件可以共享统一的依赖注入调用方式。

5.4 YAML 配置分组
配置组 当前用途
model HyperAutoModel target 和模型加载参数
tokenizer tokenizer target
training step/epoch、batch、backend、seed、梯度裁剪和 callback cadence
accelerator TP/CP/EP/PP size 与 sequence/loss parallel 标志
fsdp_config DP shard、reshard、grad sync、prefetch 等 FSDP 配置
plan_overhead 预留的 sharding plan overhead target,当前未消费
mixed_precision 预留的 mixed-precision 开关,当前训练上下文未消费
gradient_checkpointing 预留的 activation checkpoint 配置,当前 Trainer 未消费
loss_fn 可选自定义 loss module target
data_transform tokenizer 等运行时资产注入的数据变换 target
dataset transform 注入的数据集 target
collate_fn micro-batch collator target
dataloader dataset、collator 和 DP 拓扑注入的 dataloader target
optimizer model 注入的 optimizer target
lr_scheduler optimizer 和 train_steps 注入的 scheduler target
packed_sequence/magi/peft 高级数据、attention 和参数高效训练的预留配置面
checkpoint/debug/wandb 持久化、调试和远程日志的预留配置面

6. 总体设计

6.1 目标架构
train.yaml + CLI overrides
          |
          v
parse_training_args()
  YAML -> typed dataclasses + Target tree
          |
          v
TextTrainer(config)
          |
          +--> BaseTrainer._setup()
          |      process group + device + seed + DistributedSetup/MeshContext
          |
          +--> config.model.build(distributed_setup=...)
          |      HyperAutoModel -> sharding plan -> FSDP -> materialize
          |
          +--> loss / tokenizer / transform / dataset / collator / dataloader
          |
          +--> train_steps -> optimizer -> scheduler -> contexts -> callbacks
          |
          v
TextTrainer.train()
  epoch -> optimizer step -> micro batches -> forward/backward
          -> HSDP sync -> grad clip -> optimizer/scheduler -> callbacks
          |
          v
synchronize -> destroy_process_group
6.2 核心设计决策
D1:配置解析与对象构建分离

配置解析只验证结构、类型和 callable 签名,不创建占用设备内存或依赖分布式状态的对象。所有需要 model、
mesh、tokenizer、dataset 或 optimizer 的组件通过 Target.build() 延迟到 Trainer 构建阶段。

D2:分布式拓扑单一所有权

BaseTrainer._setup() 创建唯一的 DistributedSetup。AutoModel、loss 归约、dataloader 分片和 callback 指标
归约都读取同一 MeshContext,禁止各组件自行构造第二套 DeviceMesh。

D3:模型构建是原子阶段

_build_model() 调用 model target 后,模型必须已经完成本次路径需要的构造、分片、FSDP 包裹和物化。
Trainer 只在此之后构建 optimizer,避免 parameter identity、tied weights 或 sharding 改变后 optimizer 持有
过期参数引用。

D4:一个 dataloader item 对应一个 optimizer step

DataLoader 的 local batch size 为:

local_step_batch_size = global_batch_size / dp_size

每个 DP rank 的 micro-batch 数为:

num_micro_batches = global_batch_size / (micro_batch_size * dp_size)

MakeMicroBatchCollator 将 local optimizer batch 切为固定的 list[dict]。因此训练循环不需要再次猜测
梯度累积边界。

D5:训练核心显式,外围能力使用 callback

forward、backward、FSDP sync、grad clip 和 optimizer step 保持在 train_step() 中;指标、展示、评估触发
和内存清理通过 callback 生命周期接入。Callback 不修改 optimizer step 的控制流。

D6:指标只计算一次

EnvironMeterCallbackstep_train_metricsstep_env_metrics 的唯一生产者。Logging 与 Tqdm callback
只消费这两个字典,不重复做 collective 或重新计算 loss、tokens/s 和 memory 指标。


7. 构建流程

TextTrainer.__init__() 当前按下列顺序构建:

1.  _setup
2.  _build_model
3.  _build_loss
4.  _build_model_assets
5.  _build_data_transform
6.  _build_dataset
7.  _build_collate_fn
8.  _build_dataloader
9.  _compute_train_steps
10. _build_optimizer
11. _build_lr_scheduler
12. _build_training_context
13. _init_callbacks
7.1 分布式初始化

_setup() 执行:

  1. 初始化日志;
  2. 按配置 backend 初始化 torch.distributed
  3. 读取 local/global rank 和 world size;
  4. 设置当前 accelerator device;
  5. acceleratorfsdp_config 推导 DistributedSetup/MeshContext
  6. 设置随机种子和 BF16 高精度策略;
  7. rank 0 输出序列化后的 TrainerConfig;
  8. 调用当前为空实现的 save_configs()

当前主 mesh 轴顺序为:

dp_replicate? -> dp_shard? -> cp? -> tp?

demo 使用:

world_size=8
dp_shard_size=8
dp_replicate_size=1
tp=cp=ep=pp=1
mesh=(dp_shard=8)
7.2 模型与 loss

模型 target 接收 distributed_setuppeft_config。当前 demo 的 AutoModel 路径依次执行:

读取 HF config
  -> meta/no-init 构造 Qwen3-MoE
  -> 可选 ShardingPlanner
  -> per-layer + root FSDP wrapping
  -> to_empty 到当前 NPU
  -> 参数/缓冲区初始化
  -> model.train()

Loss 默认使用 ModelOutputLoss 读取 model_output.loss。若配置 loss_fn,其 target 必须构建为
torch.nn.Module,并接受 Trainer 传入的 model_outputlabels

7.3 数据链路

当前构建依赖为:

tokenizer + model_config
        -> data_transform
        -> dataset
        -> collate_fn(num_micro_batch)
        -> dataloader(dataset, collate_fn, DP topology)

demo 数据链路为:

Salesforce/wikitext / wikitext-2-raw-v1 / train
  -> PlainTextDataTransform
  -> tokenizer.encode(add_special_tokens=False)
  -> 每篇追加 EOS
  -> 串联并裁成固定 1024-token 样本
  -> input_ids + attention_mask + labels
  -> DistributedSampler
  -> MakeMicroBatchCollator
  -> list[dict]

global_batch_size 必须能被 micro_batch_size * dp_size 整除;每个 local optimizer batch 也必须能被
num_micro_batches 整除。不满足时在构建阶段报错。

7.4 训练步

每个 optimizer step 执行:

next(data_iterator) -> micro_batches
  -> global_step += 1
  -> on_step_begin
  -> device synchronize
  -> 统计整个 local step 的 loss tokens
  -> for each micro batch:
       调整 reshard_after_backward
       仅最终 micro batch 开启 FSDP gradient sync/all-reduce
       统计当前 micro batch tokens
       non_blocking 搬运到 device
       model(..., use_cache=False)
       loss_fn(model_output, labels)
       DP+CP token-weighted loss normalization
       backward
  -> hsdp_sync_stream
  -> clip_grad_norm_
  -> optimizer.step under SkipDTensorDispatch
  -> optimizer.zero_grad
  -> scheduler.step
  -> on_step_end

在读取 HSDP 异步梯度结果和执行梯度裁剪前必须完成 hsdp_sync_stream()。Optimizer 更新放在
SkipDTensorDispatch 内,避免 optimizer 对本地参数执行原始 tensor 操作时进入 DTensor dispatch。

7.5 Loss 归一化

当前路径使用 token-weighted loss:

local micro loss
  * current micro valid tokens
  / DP+CP global valid tokens for the whole optimizer step
  * dp_size

乘以 dp_size 用于抵消 FSDP 对 shard group 梯度的平均。若开启 sequence parallel,还会在 TP group
汇总 token 数并除以 sequence-parallel size。token 总数为 0 但 loss 非 0 时必须报错。

7.6 Callback 生命周期

Callback 顺序固定为:

EnvironMeterCallback
  -> LoggingCallback
  -> TqdmCallback
  -> EvaluateCallback
  -> GarbageCollectionCallback
Callback 当前行为
EnvironMeter 统计并归约 step time、tokens、samples、loss、grad norm、LR 和 device memory
Logging logging_steps 在 global rank 0 输出稳定排序的完整指标行
Tqdm global rank 0 展示从当前 global step 到 train_steps 的单一进度条
Evaluate 按 step/epoch 去重触发,但当前只记录未实现 warning
GarbageCollection 按独立 cadence 执行 Python GC 和 accelerator cache 清理

EnvironMeter 必须先于展示 callback 执行,以保证同一步的指标已经完成 collective 并写入共享字典。

7.7 训练结束与资源清理

训练达到 max_steps、epoch 结束或 dataloader 耗尽后:

  1. 调用 on_train_end()
  2. 停止 background prefetcher;
  3. accelerator synchronize;
  4. 清理 device cache;
  5. distributed barrier;
  6. 再次 synchronize;
  7. 销毁 process group。

BackgroundPrefetcher 的停止操作必须设置 stop event、清空 queue,并以有限 timeout join worker,避免退出时
无限等待后台线程。


8. 约束与失败语义

8.1 配置约束
  • modeloptimizer target 必填;
  • dataset、collate_fn 和 dataloader 在当前训练入口必填;
  • training.max_steps 若设置,必须为正整数;
  • 未设置 max_steps 时,dataloader 必须具有正的有限长度;
  • num_train_epochs 必须为正整数;
  • target 必须是可调用对象,且配置参数必须能绑定其签名;
  • 自定义 loss target 必须返回 torch.nn.Module
8.2 拓扑约束
  • world_size 必须能被 tp_size * cp_size 整除;
  • 推导出的 dp_size 必须能被 dp_shard_size 整除;
  • 当前 FSDP demo 必须存在名为 dp_shard 的 mesh 维;
  • FSDP2Manager 仅接受当前识别的 GPT-2、Llama 或 Qwen3-MoE transformer layer 结构;
  • 当前 pp_sizeep_size 尚未进入 Trainer demo 的 world-size 闭环,不得仅通过修改 YAML 视为已支持。
8.3 Batch 与 loss 约束
  • 一个 dataloader item 必须是非空 list[dict]
  • 当前 text loss 统计要求 batch 包含 labels
  • labels 中的 IGNORE_INDEX 不计入有效 token;
  • global_batch_size 必须整除 micro_batch_size * dp_size
  • model output 必须提供默认 loss 所需的 .loss,或由自定义 loss 明确处理;
  • 非阻塞搬运后的 tensor 只能在当前 device/stream 同步语义成立后读取。
8.4 Fail-closed 原则

下列情况必须明确失败,而不是静默宣称 Trainer 已支持:

  • 配置未知字段或错误类型;
  • target 路径不可导入、不是 callable 或参数签名不匹配;
  • batch size、micro batch 或拓扑无法整除;
  • FSDP 请求没有真实 DeviceMesh 或缺少 dp_shard
  • TP 被请求但 planner 没有生成有效 sharding plan;
  • model architecture 不在当前 FSDP demo 支持范围;
  • loss token 总数为 0 但模型返回非零 loss;
  • evaluation、checkpoint、PP 等占位能力被误当作完成态使用。

9. 唯一验证用例

转测qwen-moe网络,30B规格
支持tp/cp/ep/fsdp并行能力(当前迭代不支持pp),支持重计算、swap、支持梯度累加(gbs/dp/mbs > 1开启)
支持validate/production双模式,双模式精度要一致
支持断点续训,集群故障导致训练中断,拉起后可恢复至断点状态继续训练。
支持在线加载HF数据集;离线转换hf数据集为megatron格式数据集,然后加载megatron格式数据集

转测需要验收:
1、精度的自洽,保证输入、权重一致的情况下,变换并行策略、重计算策略、梯度累加长度等保持loss曲线在500step下符合验收标准(0轴波动,误差范围在xxx以内)

9.1 examples/training_demo 端到端 Trainer 跑通

本 RFC 只定义这一条验证用例,不增加单元测试、组合并行矩阵、数值对拍或性能测试。

环境要求:

  • 8 张可用 Ascend NPU;
  • PyTorch、torch-npu、HCCL 和 HyperParallel 依赖可用;
  • 能访问或已缓存 Qwen/Qwen3-30B-A3B config/tokenizer;
  • 能访问或已缓存 Salesforce/wikitext
  • host 内存、磁盘和 NPU HBM 足以完成随机 checkpoint 准备与 8 卡 FSDP 训练。

执行命令:

bash examples/training_demo/run.sh

该脚本串行执行:

python -m examples.training_demo.prepare_model examples/training_demo/train.yaml
  -> torchrun --nproc_per_node=8 --module examples.training_demo.train_text \
       examples/training_demo/train.yaml

默认配置:

项目
model Qwen3-30B-A3B,BF16,SDPA,HF path
dataset WikiText-2 raw train
sequence length 1024
max steps 100
global/micro batch 8 / 1
backend HCCL
topology 8-way FSDP,TP=CP=EP=PP=1
optimizer AdamW
scheduler 1-step warmup + cosine decay
checkpoint disabled

10. 验收标准

  • 模型、tokenizer、dataset、collator、dataloader、optimizer、scheduler 和 callbacks 均成功构建;
  • 日志确认真实 (dp_shard=8) DeviceMesh 和 Qwen3-MoE per-layer/root FSDP wrapping;
  • 8 个 rank 完成 100 个 optimizer steps,无 HCCL hang、未处理异常或进程提前退出;
  • 每步 loss、grad norm 和 learning rate 为有限值;
  • global rank 0 的日志或 tqdm 能持续展示 step、loss、grad norm、LR、step time 和 tokens/s;
  • 最终执行 accelerator synchronize 和 process-group 销毁,命令退出码为 0。

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

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 hyper_models/trainer/config.py, base.py, and text_trainer.py, then inspect examples/training_demo/train.yaml and train_text.py. Run the documented CLI entry point in the target 8-card Ascend setup and compare the behavior with the RFC's support matrix. Done means the documented Trainer responsibilities and end-to-end Qwen3-30B-A3B/WikiText demo agree with the implemented scope.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
distributed-systems, machine-learning
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Active
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.