mindspore-ai / mindspore-ai/hyper-parallel
[RFC] Torch Qwen3-30B-A3B 重计算
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 53
- Forks
- 63
- Avg merge
- 23h 45m
- Merged PRs (30d)
- 63
Description
1. 基本信息
| 项目 | 内容 |
|---|---|
| 作者 | 宋佳琪 |
| 相关模块 | model / activation_checkpoint / trainer |
| 相关 issue / PR | https://gitcode.com/mindspore/hyper-parallel/pull/1132 |
| 适用后端 | PyTorch(运行时设置 HYPER_PARALLEL_PLATFORM=torch) |
| 配置入口 | examples/training_demo/train.yaml 中的 activation_checkpoint.mode |
| 适用模型 | 本期仅验收只承诺 Qwen3-30B-A3B |
2. 背景
大模型训练的前向过程会为反向传播保留大量中间激活。随着模型层数、序列长度和 micro batch 增大,激活可能成为设备峰值显存的主要来源,并限制可训练的模型规模。
Activation Checkpoint(重计算)通过在前向阶段少保存部分激活、在反向阶段重新执行对应前向计算,以额外计算量换取设备显存。Hyper-Parallel 已提供通用的 checkpoint_wrapper 和选择性重计算能力;Trainer 在模型准备阶段增加统一配置入口,负责识别 Hugging Face 模型中的 Transformer layer,适配重计算。
本功能解决的问题:用户只需在 Trainer YAML 的 activation_checkpoint.mode配置 off、full 或 selective,即可为 Trainer 拉起的模型启用整层选择性重计算或子模块完全重计算。
成功标准:开启后 loss 和梯度与关闭重计算的基线满足项目精度要求,反向阶段确实发生目标区域的重新计算,且目标训练场景的设备峰值显存下降。
3. 目标和非目标
3.1 目标
- 仅支持 PyTorch 后端,验收模型为 Qwen3-30B-A3B。
- 在 Trainer 公共配置中提供
off、full、selective三种模式,默认关闭,保持现有训练行为。
3.2 非目标
- 不支持 MindSpore。
- YAML 不开放自定义
policy_fn、swap_inputs等底层参数;selective使用 Trainer 内置策略。 - 不支持与
activation_swap=attention同时开启,两者同时配置时在模型准备阶段报错。
4. 相关实现参考
| 来源 | 做法 | 对 Trainer 适配的影响 |
|---|---|---|
Hyper-Parallel checkpoint_wrapper |
使用非可重入 checkpoint 包装 module,反向按需重放 forward | 作为整层和子模块重计算的统一包装能力 |
Hyper-Parallel CheckpointPolicy |
按算子返回 MUST_SAVE / MUST_RECOMPUTE |
用于 selective 固定策略 |
| Hugging Face gradient checkpointing | 模型原生识别 GradientCheckpointingLayer 并管理 checkpoint 调用 |
满足条件时作为 full 的优先路径 |
| PyTorch selective checkpoint context | 在同一 checkpoint 区域记录 forward/recompute 算子序列 | 支持按算子保存或重算,并要求两次执行可一致重放 |
5. 对外接口
5.1 接口定义
Trainer 配置定义:
@dataclass
class ActivationCheckpointConfig:
mode: Optional[Literal["off", "full", "selective"]] = "off"
YAML 示例:
activation_checkpoint:
mode: selective
| 配置项 | 类型 | 默认值 | 是否必填 | 含义 | 合法范围 | 错误处理 |
|---|---|---|---|---|---|---|
activation_checkpoint.mode |
str |
off |
否 | Trainer 重计算模式 | off / full / selective |
非法值在配置解析阶段报错 |
5.2 模式语义
| 模式 | 行为 | checkpoint 粒度 |
|---|---|---|
off |
不调用 Trainer 重计算适配 | 无 |
full |
优先使用 HF 原生重计算;不满足条件时用 checkpoint_wrapper 包装 layer 内的计算子模块 | HF layer 或 attention/MLP/norm 子模块 |
selective |
checkpoint_wrapper 包装完整 Transformer layer,并按内置算子策略保存昂贵结果、重算其余结果 | 完整 Transformer layer 内的算子 |
5.3 使用示例
关闭重计算:
activation_checkpoint:
mode: off
开启 full 重计算:
activation_checkpoint:
mode: full
开启 selective 重计算:
activation_checkpoint:
mode: selective
5.4 参数校验
启用 full 或 selective 时需要满足:
- 模型能够解析出至少一个非空 Transformer layer 容器,否则抛出
ValueError。 activation_swap必须为none,否则抛出不兼容错误。
6. 方案设计
6.1 Transformer layer 容器识别
容器发现以 nn.ModuleList 或数字 key 的 nn.ModuleDict 为边界。识别顺序如下:
| 模型类型 | 识别方式 | 约束 |
|---|---|---|
| 已登记多模态/语言模型 | 按模型类名查找预定义的 language/vision 路径,每个 role 只取第一个有效路径 | 已登记模型不使用未知模型启发式兜底 |
| 常见未登记语言模型 | 依次查找 model.layers、layers |
容器必须非空 |
| 未知模型 | 查找最大的 ModuleList 或数字 key ModuleDict |
仅为保守启发式,可能无法表达多 tower 结构 |
| Retrieval wrapper | 对内部 model 递归识别 |
当前识别 BiEncoderModel、CrossEncoderModel、FSDPBiEncoderModel |
当前显式登记覆盖 Gemma3/Gemma4、Qwen2-VL/Qwen2.5-VL、Qwen3.5/Qwen3-VL、LLaVA、Mistral3/Ministral3、Llama4/Llama-Nemotron-VL、SmolVLM、Kimi-VL、MiniMax-M3、Step3.7、Bagel、Nemotron-H 和 GPT-2 等结构。登记表示 Trainer 知道 layer 容器位置,不等同于所有模型和并行组合都已完成端到端验收。
数字 key ModuleDict 用于兼容 pipeline split 后只保留本 rank layer 的结构;非数字 key 的 ModuleDict 不会被未知模型启发式误认为 Transformer layer 容器。
6.2 full 模式
full 按以下优先级选择实现:
flowchart TD
A[full] --> B{仅 language tower}
B -- 否 --> F[HP 子模块 checkpoint]
B -- 是 --> C{未开启 compile}
C -- 否 --> G[compile 专用子模块 checkpoint]
C -- 是 --> D{所有 layer 可训练且继承 GradientCheckpointingLayer}
D -- 否 --> F
D -- 是 --> E{模型声明 supports_gradient_checkpointing 且提供 enable API}
E -- 是 --> H[HF native use_reentrant=True]
E -- 否 --> F
| 场景 | 实际实现 | 包装范围 |
|---|---|---|
| 满足 HF 原生条件 | gradient_checkpointing_enable(...use_reentrant=True) |
Hugging Face 自己管理 language layer 重计算 |
| 不满足 HF 原生条件且未 compile | Hyper-Parallel 非可重入 checkpoint_wrapper |
每层已知名称的 MLP、Attention、两处 norm,以及存在时的 MTP/MoE generation 子模块 |
开启 torch.compile |
Hyper-Parallel 非可重入 checkpoint_wrapper |
每层已知名称的 Attention 和 MLP/FFN 子模块,不包装 norm |
常规子模块识别别名:
MLP: mlp / feed_forward / ffn
Attention: self_attn / attention / attn / linear_attn
Norm 1: input_layernorm / attention_norm / layer_norm1 / norm1
Norm 2: post_attention_layernorm / ffn_norm / layer_norm2 / norm2
MTP: mlp_moe_gen / input_layernorm_moe_gen / post_attention_layernorm_moe_gen
full 并不表示无条件包装整个 decoder layer。HF 原生路径的粒度由模型实现决定;Hyper-Parallel fallback 则只包装上述可识别子模块。若自定义 layer 不使用这些名称,可能出现包装数量为 0,需通过日志和实际 backward 调用次数确认功能是否生效。
6.3 selective 模式
无 KV 共享时,selective 使用非可重入 checkpoint_wrapper 包装发现到的每个完整 Transformer layer,并为每个 checkpoint invocation 创建独立的算子计数和 forward/recompute context。
内置策略如下:
| 算子类别 | 策略 | 原因 |
|---|---|---|
topk |
MUST_RECOMPUTE |
部分模型会原地修改其输出,保存引用会触发版本检查或重放已修改 tensor |
matmul / mm / linear / grouped matmul |
同一区域内按出现次序交替 MUST_SAVE、MUST_RECOMPUTE |
在矩阵计算开销和激活显存之间折中 |
| 已知昂贵计算、Attention kernel 和通信算子 | MUST_SAVE |
避免在 backward 重复高开销计算或集合通信 |
| 其他算子 | MUST_RECOMPUTE |
释放普通中间激活 |
昂贵算子集合包含当前 PyTorch 可提供的 compute-intensive op、GEMM/BMM、SDPA/Flash/Flex/FFPA/torch-attn、NPU fusion attention、EP dispatch/combine、all-to-all、reduce-scatter 和 all-reduce 等。可选算子只在当前运行环境已注册时加入策略。
Profiler record-function op 和 FSDP 参数生命周期相关的 allocation/copy/all-gather op 会从 selective replay 计数中忽略。这些操作仍正常执行,但不会因为 forward prefetch 与 recompute 时机不同而破坏 SAC 算子序列匹配。
6.4 KV 共享和 cache 处理
Trainer 通过 model.config.text_config.num_kv_shared_layers > 0 判断跨层 KV 共享:
| 场景 | cache 行为 | 重计算行为 |
|---|---|---|
| 无 KV 共享 | 尽力将主 config 及其 sub-config 的 use_cache 设为 False |
按所选 full 或 selective 路径执行 |
| 有 KV 共享 | 保留 use_cache |
跳过 Attention,避免 backward 重算再次写共享 KV cache |
KV 共享模型选择 selective 时会记录 warning,并降级为子模块 checkpoint:包装 MLP 和 norm,跳过 Attention。此时不再是完整 layer 的 selective 算子策略,显存收益也可能低于普通 selective。
6.5 代码改动点
| 模块 | 职责 | 默认是否影响已有行为 |
|---|---|---|
hyper_models/trainer/config.py |
定义 ActivationCheckpointConfig.mode |
默认 off,不影响 |
hyper_models/trainer/base.py |
在 _build_model() 中向 model target 透传 mode |
off 时只透传配置 |
hyper_models/_transformers/auto_model.py |
接收 mode 并传入模型 infrastructure | off 时不包装 |
hyper_models/_transformers/infrastructure.py |
控制 sharding、重计算、compile、FSDP2 的执行顺序 | 只在显式开启时调用适配 |
hyper_models/components/distributed/activation_checkpointing.py |
容器发现、full 路由、selective policy、KV 共享处理 | 只在显式开启时执行 |
hyper_parallel/core/activation_checkpoint |
提供 checkpoint wrapper、policy 和 recompute context | Trainer 直接复用,不改变公共 API |
7. 组件依赖
| 依赖组件 | 强依赖 / 弱依赖 | 当前状态 | 未 ready 时能力 |
|---|---|---|---|
| PyTorch autograd/checkpoint | 强依赖 | 已有 | 无法提供 Trainer 重计算 |
Hyper-Parallel checkpoint_wrapper |
强依赖 | 已有 | full fallback 和 selective 无法工作 |
| Hugging Face Transformers | 强依赖 | 已有 | 无法构建 Trainer 目标模型和使用 HF 原生路径 |
| FSDP2 | 弱依赖 | 已适配执行顺序和 selective ignore op | 可在无 FSDP 场景使用;正式组合需单独验证 |
| TP/CP/EP sharding | 弱依赖 | 在重计算前应用 | 可单独使用重计算;组合需按目标模型验证精度和通信次数 |
torch.compile |
弱依赖 | full 有专用子模块路径,selective 按 wrapper 后 compile 执行 |
不开启 compile 不影响重计算基本能力 |
| Activation Attention Swap | 互斥依赖 | 当前不支持同时开启 | 同时配置时快速失败 |
| MindSpore | 不涉及 | Trainer 适配未实现 | 不支持 |
最小可交付能力:PyTorch + Hugging Face Qwen3-30B-A3B 通过 YAML 开启 full 或 selective,完成训练并观察到重计算。
8. 约束与兼容性
| 类型 | 内容 |
|---|---|
| 后端 | Trainer 适配仅支持 PyTorch |
| 默认兼容性 | 默认 mode=off ,不调用重计算 |
| Activation swap | activation_swap=attention 与任何非 off 重计算互斥,模型准备阶段报错 |
| Checkpoint 文件 | state_dict key 必须与未包装模型兼容;save/load 后可继续训练 |
| 显存收益 | 取决于激活占比和实际保存策略;KV 共享降级、短序列或小模型可能收益有限 |
| 性能代价 | backward 增加 forward 计算;selective 保存昂贵计算结果以降低代价,但不保证达到某固定吞吐 |
| 版本依赖 | 模型 class 名、module tree、可选 op 注册和 PyTorch 私有 compute-intensive op 列表均可能随 Transformers/PyTorch 版本变化 |
9. 验证设计
9.1 用例分层
| 用例级别 | 覆盖内容 | 通过标准 |
|---|---|---|
| UT | 配置解析、关闭模式无副作用、full/selective 包装与 backward 重算、KV 共享、模型容器识别、HF 原生路由、compile 路由、sharding/FSDP 顺序 | 全部通过,目标分支和失败路径可观测 |
| Level0 | tiny causal LM 的 off/full/selective 单进程对拍 |
loss、输入梯度和参数梯度满足精度阈值;重算调用次数符合预期 |
| Level1 | 目标大模型的 FSDP2/TP/CP/EP 实际训练组合 | 连续多 step 无 hang/OOM,loss/梯度符合项目标准,峰值显存低于基线 |
| 兼容性 | HF 原生、VLM 多 tower、KV 共享、compile、checkpoint save/load/resume | 执行路径符合设计,恢复训练后结果连续 |
9.2 核心正确性验证
- 固定随机种子、输入、dtype、优化器和并行拓扑,对比
off、full、selective的 forward loss 和 parameter gradients。 - 证明 backward 确实重放目标 layer/submodule,而不是只完成 wrapper 替换。
9.3 交互验证
| 组合 | 是否支持 | 通过标准 |
|---|---|---|
full/selective + FSDP2/TP/CP/EP |
是,按目标组合验收 | loss/梯度对齐,无通信错误 |
full/selective + torch.compile |
是 | 配置允许,forward/backward 和基线对齐 |
| 重计算 + Attention swap | 否 | 模型准备阶段按预期报错 |
| 重计算 + MindSpore Trainer | 否 | 不作为本功能支持场景 |
9.4 性能和显存验证
性能和显存不预设固定收益比例。
Qwen3-30B-A3B seq_length:1024
| 配置 | peak_mem | 显存优化率 | step_time | 性能劣化率 |
|---|---|---|---|---|
off |
29.60G | / | 6.2919s | / |
full |
26.91G | 9.1% | 7.2091s | 14.58% |
selective |
25.74G | 13% | 7.9122s | 25.88% |
10. 验收 Checklist
- 配置默认值为
off,不使能重计算。 - 目标大模型配置
off/full/selective时的 loss 和梯度与基线对齐。 - 重计算与 FSDP2/TP/CP/EP/compile 组合完成连续训练。
- Attention swap 与重计算同时开启时快速失败。
- 目标场景的峰值显存下降。
schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 336
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/336
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, hyper_models/trainer/base.py, hyper_models/_transformers/auto_model.py, hyper_models/_transformers/infrastructure.py, and hyper_models/components/distributed/activation_checkpointing.py. Review the UT and tiny causal LM validation layers, then verify off/full/selective behavior, gradient and loss alignment, incompatibility errors, recomputation, and peak-memory reduction for Qwen3-30B-A3B.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python, pytorch
- Domain
- distributed-systems, machine-learning, performance
- Issue type
- Feature
- Difficulty
- 5/5
- Estimated time
- Over a week
- Activity status
- Stale
- Clarity
- Mostly clear
- Newbie friendliness
- 25/100