mindspore-ai / mindspore-ai/hyper-parallel
[RFC] end2end trainer w/ Qwen3-30B training demo
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 以及多种组合并行能力,但完整训练
还需要解决另一组编排问题:
- 如何把 YAML 配置解析为有类型、可校验的运行时组件;
- 如何只创建一次分布式拓扑,并让模型、数据、loss 和指标共享同一语义;
- 如何保证模型完成构建、并行化和参数物化后,再创建 optimizer 等依赖参数身份的组件;
- 如何把 global batch 拆成每个 DP rank 的 optimizer batch 和多个 micro batch;
- 如何在梯度累积下正确控制 FSDP reshard、gradient sync、梯度裁剪和 optimizer step;
- 如何把日志、进度、环境指标、评估触发和内存清理从核心训练步中拆出;
- 如何提供一个真实模型、真实数据和多卡 FSDP 的端到端入口,证明上述组件能够串成闭环。
当前分支已经形成以 TrainerConfig + Target + BaseTrainer + TextTrainer + Callback 为核心的实现,并新增
examples/training_demo 作为 Qwen3-30B-A3B/WikiText/FSDP 示例。需要通过 RFC 固化当前职责边界,明确
已实现能力与占位接口,避免把配置面的可表达性误认为运行时支持。
2. 目标与非目标
2.1 目标
本 RFC 固化以下能力:
- 使用一个 YAML 文件描述模型、tokenizer、数据变换、dataset、collator、dataloader、loss、optimizer、
scheduler、并行拓扑和 callback cadence。 - 通过
parse_training_args()解析配置,并支持--section.field=value形式的强类型 CLI 覆盖。 - 通过
Target延迟构建需要运行时依赖的组件,配置解析阶段不提前创建 model、dataset 或 optimizer。 TextTrainer作为当前文本训练规范入口,按显式依赖顺序组装BaseTrainer的各阶段。- 分布式初始化和
DistributedSetup/MeshContext只创建一次,并注入 AutoModel 构建路径。 - 模型由
HyperAutoModelForCausalLM.from_pretrained()原子完成构造、分片、FSDP 包裹和物化;Trainer
不再执行第二次并行化。 - DataLoader 每次迭代返回
list[dict[str, Any]],该列表表示一个 optimizer step 内的全部 micro batches。 - 训练步完成 token 统计、forward/backward、FSDP 累积同步控制、梯度裁剪、optimizer step、scheduler
step 和指标发布。 - 使用 callback 承担环境指标、结构化日志、tqdm、评估触发占位和周期性内存清理。
- 使用唯一的
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_steps、empty_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 不接受 **kwargs,Target.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:指标只计算一次
EnvironMeterCallback 是 step_train_metrics 和 step_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() 执行:
- 初始化日志;
- 按配置 backend 初始化
torch.distributed; - 读取 local/global rank 和 world size;
- 设置当前 accelerator device;
- 由
accelerator与fsdp_config推导DistributedSetup/MeshContext; - 设置随机种子和 BF16 高精度策略;
- rank 0 输出序列化后的 TrainerConfig;
- 调用当前为空实现的
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_setup 和 peft_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_output 与 labels。
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 耗尽后:
- 调用
on_train_end(); - 停止 background prefetcher;
- accelerator synchronize;
- 清理 device cache;
- distributed barrier;
- 再次 synchronize;
- 销毁 process group。
BackgroundPrefetcher 的停止操作必须设置 stop event、清空 queue,并以有限 timeout join worker,避免退出时
无限等待后台线程。
8. 约束与失败语义
8.1 配置约束
model与optimizertarget 必填;- 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_size和ep_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-A3Bconfig/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
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- 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