已开启
[RFC] Torch Qwen3-30B-A3B 重计算 #336
songjiaqi创建于 8月17日
8月17日 添加了label:RFC
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 将 silkage_jiajia 设为负责人
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月18日 修改了issue 的描述
1. 基本信息
model/activation_checkpoint/trainerHYPER_PARALLEL_PLATFORM=torch)examples/training_demo/train.yaml中的activation_checkpoint.mode2. 背景
大模型训练的前向过程会为反向传播保留大量中间激活。随着模型层数、序列长度和 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 目标
off、full、selective三种模式,默认关闭,保持现有训练行为。3.2 非目标
policy_fn、swap_inputs等底层参数;selective使用 Trainer 内置策略。activation_swap=attention同时开启,两者同时配置时在模型准备阶段报错。4. 相关实现参考
checkpoint_wrapperCheckpointPolicyMUST_SAVE/MUST_RECOMPUTEselective固定策略GradientCheckpointingLayer并管理 checkpoint 调用full的优先路径5. 对外接口
5.1 接口定义
Trainer 配置定义:
@dataclass class ActivationCheckpointConfig: mode: Optional[Literal["off", "full", "selective"]] = "off"YAML 示例:
activation_checkpoint: mode: selectiveactivation_checkpoint.modestroffoff/full/selective5.2 模式语义
offfullselective5.3 使用示例
关闭重计算:
activation_checkpoint: mode: off开启 full 重计算:
activation_checkpoint: mode: full开启 selective 重计算:
activation_checkpoint: mode: selective5.4 参数校验
启用
full或selective时需要满足:ValueError。activation_swap必须为none,否则抛出不兼容错误。6. 方案设计
6.1 Transformer layer 容器识别
容器发现以
nn.ModuleList或数字 key 的nn.ModuleDict为边界。识别顺序如下:model.layers、layersModuleList或数字 keyModuleDictmodel递归识别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 -- 否 --> Fgradient_checkpointing_enable(...use_reentrant=True)checkpoint_wrappertorch.compilecheckpoint_wrapper常规子模块识别别名:
full并不表示无条件包装整个 decoder layer。HF 原生路径的粒度由模型实现决定;Hyper-Parallel fallback 则只包装上述可识别子模块。若自定义 layer 不使用这些名称,可能出现包装数量为 0,需通过日志和实际 backward 调用次数确认功能是否生效。6.3
selective模式无 KV 共享时,
selective使用非可重入checkpoint_wrapper包装发现到的每个完整 Transformer layer,并为每个 checkpoint invocation 创建独立的算子计数和 forward/recompute context。内置策略如下:
topkMUST_RECOMPUTEmatmul/mm/linear/ grouped matmulMUST_SAVE、MUST_RECOMPUTEMUST_SAVEMUST_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 共享:use_cache设为Falsefull或selective路径执行use_cacheKV 共享模型选择
selective时会记录 warning,并降级为子模块 checkpoint:包装 MLP 和 norm,跳过 Attention。此时不再是完整 layer 的 selective 算子策略,显存收益也可能低于普通selective。6.5 代码改动点
hyper_models/trainer/config.pyActivationCheckpointConfig.modeoff,不影响hyper_models/trainer/base.py_build_model()中向 model target 透传 modeoff时只透传配置hyper_models/_transformers/auto_model.pyoff时不包装hyper_models/_transformers/infrastructure.pyhyper_models/components/distributed/activation_checkpointing.pyhyper_parallel/core/activation_checkpoint7. 组件依赖
checkpoint_wrapperfullfallback 和selective无法工作torch.compilefull有专用子模块路径,selective按 wrapper 后 compile 执行最小可交付能力:PyTorch + Hugging Face Qwen3-30B-A3B 通过 YAML 开启
full或selective,完成训练并观察到重计算。8. 约束与兼容性
mode=off,不调用重计算activation_swap=attention与任何非off重计算互斥,模型准备阶段报错selective保存昂贵计算结果以降低代价,但不保证达到某固定吞吐9. 验证设计
9.1 用例分层
off/full/selective单进程对拍9.2 核心正确性验证
off、full、selective的 forward loss 和 parameter gradients。9.3 交互验证
full/selective+ FSDP2/TP/CP/EPfull/selective+torch.compile9.4 性能和显存验证
性能和显存不预设固定收益比例。
Qwen3-30B-A3B seq_length:1024
offfullselective10. 验收 Checklist