已开启
[RFC] Torch Qwen3-30B-A3B 重计算 #336
songjiaqi创建于  8月17日
songjiaqi
songjiaqi成员
8月17日 创建

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配置 offfullselective,即可为 Trainer 拉起的模型启用整层选择性重计算或子模块完全重计算。

成功标准:开启后 loss 和梯度与关闭重计算的基线满足项目精度要求,反向阶段确实发生目标区域的重新计算,且目标训练场景的设备峰值显存下降。

3. 目标和非目标

3.1 目标

  1. 仅支持 PyTorch 后端,验收模型为 Qwen3-30B-A3B。
  2. 在 Trainer 公共配置中提供 offfullselective 三种模式,默认关闭,保持现有训练行为。

3.2 非目标

  1. 不支持 MindSpore。
  2. YAML 不开放自定义 policy_fnswap_inputs 等底层参数;selective 使用 Trainer 内置策略。
  3. 不支持与 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 参数校验

启用 fullselective 时需要满足:

  1. 模型能够解析出至少一个非空 Transformer layer 容器,否则抛出 ValueError
  2. activation_swap 必须为 none,否则抛出不兼容错误。

6. 方案设计

6.1 Transformer layer 容器识别

容器发现以 nn.ModuleList 或数字 key 的 nn.ModuleDict 为边界。识别顺序如下:

模型类型 识别方式 约束
已登记多模态/语言模型 按模型类名查找预定义的 language/vision 路径,每个 role 只取第一个有效路径 已登记模型不使用未知模型启发式兜底
常见未登记语言模型 依次查找 model.layerslayers 容器必须非空
未知模型 查找最大的 ModuleList 或数字 key ModuleDict 仅为保守启发式,可能无法表达多 tower 结构
Retrieval wrapper 对内部 model 递归识别 当前识别 BiEncoderModelCrossEncoderModelFSDPBiEncoderModel

当前显式登记覆盖 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_SAVEMUST_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 按所选 fullselective 路径执行
有 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 开启 fullselective,完成训练并观察到重计算。

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 核心正确性验证

  1. 固定随机种子、输入、dtype、优化器和并行拓扑,对比 offfullselective 的 forward loss 和 parameter gradients。
  2. 证明 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

likedislike
songjiaqisongjiaqi成员
8月17日 添加了label:RFC
songjiaqisongjiaqi成员
8月17日 修改了issue 的描述
songjiaqisongjiaqi成员
8月17日 修改了issue 的描述
songjiaqisongjiaqi成员
8月17日 修改了issue 的描述
songjiaqisongjiaqi成员
8月17日 修改了issue 的描述
songjiaqisongjiaqi成员
8月17日 修改了issue 的描述
songjiaqisongjiaqi成员
8月17日 修改了issue 的描述
songjiaqisongjiaqi成员
8月17日 修改了issue 的描述
songjiaqisongjiaqi成员
8月17日 修改了issue 的描述
songjiaqisongjiaqi成员
8月17日 将 silkage_jiajia 设为负责人
songjiaqisongjiaqi成员
8月17日 修改了issue 的描述
songjiaqisongjiaqi成员
8月17日 修改了issue 的描述
songjiaqisongjiaqi成员
8月18日 修改了issue 的描述