mhc-recompute-fboverlap
recompute-g2-attention
--recompute-csa-attention
--mhc-recompute
MindSpeed-LLM 分支名沿用开发初期名称,其中的 g2-attention 仅是历史分支标识。当前代码、接口和文档统一使用 CSA 命名。
g2-attention
本提案为 DeepSeek V4 Mcore 训练增加细粒度重计算能力。方案不对整个 Transformer Layer 做统一重计算,而是选择 CSA 与融合 MHC 中内存占用较高、计算代价可接受的中间结果,在前向阶段主动释放其底层存储,并在反向传播即将使用时恢复,从而降低训练峰值显存。
方案由两个仓库协同实现:Megatron/MindSpeed 提供无输出保存重计算基础能力,MindSpeed 提供 Pipeline Parallel(PP)层优先级计算以及 MoE forward-backward overlap(以下简称 FB overlap)中的释放时序;MindSpeed-LLM 提供 DeepSeek V4 模型侧的 CSA indexer、稀疏注意力、输出投影、compressor 和 MHC 重计算逻辑。用户通过 --recompute-csa-attention 和 --mhc-recompute 两个参数控制主要能力。
DeepSeek V4 的 CSA 前向链路包含 q-up projection、q-norm/RoPE、indexer、压缩 KV、稀疏注意力、o-down projection 和 o-up projection;启用 MHC 后还会增加 MHC pre/post 的中间结果。在长序列、MoE、FB overlap、PP/VPP 或自定义 pipeline-model-parallel-layout 组合下,这些激活可能在反向传播前长期驻留,成为峰值显存的重要组成部分。
pipeline-model-parallel-layout
当前实现需要同时处理多个子模块之间的生命周期依赖:CSA o-up 输出会继续进入 BDA 和 MHC post;FB overlap 会拆分并交错执行前向、反向与通信;MTP 和自定义 PP layout 还会改变层构建与 layer/chunk 映射。因此,细粒度重计算不仅需要划分合适的重计算子图,还需要在普通 Transformer 与 FB overlap 路径中选择正确的输出释放点和反向挂钩点。
本方案使用 --recompute-csa-attention 统一控制 CSA 子图,在内部对 q-norm、sparse attention 和 o-up 分别管理 checkpoint 生命周期,并通过 MindSpeed 的 overlap 调度协调 BDA、MHC 和 norm 重计算。对外仅注册功能级参数,不暴露子图级开关。
用户希望在不启用整层重计算的情况下,仅重计算 CSA 中显存占用较高的子图。开启 --recompute-csa-attention 后,方案覆盖:
--recompute-norm
关键要求是训练 loss 与不开启细粒度重计算的基线在既有数值容差内对齐,且反向梯度图完整。
当同时使用 --enable-mhc --use-fused-mhc --mhc-recompute --recompute-csa-attention 时,CSA 的输出会进入 attention BDA 和 MHC post。释放点必须同时满足:
--enable-mhc --use-fused-mhc --mhc-recompute --recompute-csa-attention
(y, post, comb)
FB overlap 会拆分并交错执行不同层的前向和反向图。方案需要在 MindSpeed 的 overlap 调度中保存 checkpoint manager,并在合适的 layer output 上注册重计算 hook,同时不干扰 All-to-All、专家计算、权重梯度解耦和既有 swap 流程。
普通 MoE FB overlap 与 balanced MoE 使用不同的 attention_forward 引用,因此 MindSpeed-LLM 需要通过 feature patch 将统一的 CSA 释放包装器安装到对应入口。
attention_forward
细粒度重计算通常会与 recompute-num-layers、recompute-norm-num-layers 等按层配置一起使用。自定义 PP layout 下,每个 PP/VPP chunk 的 decoder layer 数可能不均匀,不能再用全局平均层数推导 chunk 大小。
recompute-num-layers
recompute-norm-num-layers
方案要求直接解析 pipeline_model_parallel_layout,依据当前 PP rank、VPP rank 和 decoder layer 的实际位置计算重计算优先级。MTP 层单独标记并跳过当前 CSA/MHC 细粒度重计算,以避免把主模型层的生命周期假设错误应用到 MTP 图。
pipeline_model_parallel_layout
细粒度重计算与 --swap-layer-input 可以组合使用,但两者职责不同:
--swap-layer-input
FB overlap 会使不同层、不同 microbatch 的前反向交错。为避免仅根据相邻层和 manager 队列顺序恢复到错误输入,当前实现将每次换出产生的 swap_entry 与对应 LayerGraph 绑定,在 bwd_layer_graph 或 next_bwd_layer_graph 进入反向前精确恢复。恢复完成后按对象身份从 manager 中移除 entry,避免扰乱其他 microbatch 的队列。
swap_entry
LayerGraph
bwd_layer_graph
next_bwd_layer_graph
MindSpeed 保留面向 Megatron TransformerLayer 的默认 feature,同时提供通用 swap wrapper 和 manager;MindSpeed-LLM 保留自己的 SwapLayerInputFeature,负责向自定义 TransformerLayer 及所需 FB overlap 入口注册 patch。MindSpeed-LLM 不重复实现 swap wrapper 和 manager。MTP Layer 不参与当前 layer-input swap manager 调度,避免重复构建影响主模型的 manager 数量和层编号。
SwapLayerInputFeature
False
training
torch.is_grad_enabled()
整体设计分为模型策略层、通用重计算层和 overlap 生命周期层。
flowchart LR A["CLI 参数"] --> B["MindSpeed-LLM Feature Manager"] B --> C["DeepSeek V4 CSA 与 MHC 子图"] C --> D["CheckpointWithoutOutput"] D --> E["前向计算并保留 autograd 上下文"] E --> F["到达安全释放点"] F --> G["释放输出 storage 并注册 backward hook"] G --> H["反向到达 hook"] H --> I["恢复 RNG 并重算输出"] I --> J["继续原反向传播"] K["MindSpeed FB overlap 调度"] --> F L["PP layout 重计算优先级"] --> C M["Swap layer input"] --> N["LayerGraph 绑定 swap entry"] N --> O["反向前换回 Layer 输入"]
职责划分如下:
recompute_priority
当前代码采用以下技术组件:
CheckpointWithoutOutput
is_mtp_layer
is_mtp_attention
--recompute-csa-attention 开启后,DeepSeek V4 CSA 依据运行条件建立以下 checkpoint:
kv_compress
重计算只在以下条件同时满足时启用:
参数开启 AND 非 MTP Attention AND module.training AND torch.is_grad_enabled()
CSA 与 MHC post 同时重计算时,生命周期分为两个阶段:
sequenceDiagram participant CSA as CSA participant BDA as Attention BDA participant MHC as MHC Post participant L as Layer/FB Overlap participant BW as Backward CSA->>CSA: 计算 q-norm、sparse、o-up CSA->>BDA: 返回 o-up 输出 BDA->>MHC: 生成 BDA 输出并进入 MHC post MHC-->>L: 返回 MHC post 输出 L->>CSA: 立即释放 q-norm/sparse 中间输出 Note over CSA,L: o-up 暂时保留 L->>L: 继续 MLP/MHC,形成稳定 hook tensor L->>CSA: 释放 o-up,并为所有相关 checkpoint 注册 hook L->>MHC: 释放受支持的 MHC pre/post 输出并注册 hook BW->>CSA: hook 触发,按依赖顺序重算 BW->>MHC: 恢复 MHC 输出并继续反向
MindSpeed-LLM 的 discard_csa_attention_intermediate_outputs() 只清理 q-norm 和 sparse checkpoint 输出;discard_csa_attention_output(hook_tensor) 负责最终挂接 q-norm、sparse 与 o-up 的重计算。csa_fb_overlap_attention_forward_wrapper() 根据 defer_attention_recompute_for_mhc_post 判断采用单阶段还是分阶段释放。
discard_csa_attention_intermediate_outputs()
discard_csa_attention_output(hook_tensor)
csa_fb_overlap_attention_forward_wrapper()
defer_attention_recompute_for_mhc_post
该设计解决两类相反问题:
--mhc-recompute 面向融合 NPU MHC 实现,当前覆盖 attention/MLP MHC pre 以及 attention MHC post:
x、residual、post、comb
is_mtp_layer=True
当 CSA 与 MHC post 需要统一延迟挂钩时,MindSpeed 将 attention BDA 也纳入 CheckpointWithoutOutput。根据 attention bias 是否存在,分别构造无 bias 和有 bias 的 checkpoint 函数,保持 BDA 调用签名稳定。
--recompute-norm 开启时:
defer_attention_recompute_norm
MindSpeed 的 FB overlap forward/backward 图增加以下能力:
MindSpeed-LLM 仅在以下条件成立时注册 FB overlap patch:
position_embedding_type == deepseek4 AND moe_fb_overlap == true AND recompute_csa_attention == true
balanced MoE 开启时额外 patch balanced MoE 的 forward 与 forward-backward 入口。包装器带 _csa_fb_overlap_wrapped 标识,防止重复包装。
_csa_fb_overlap_wrapped
MindSpeed 提供统一的 get_recompute_priority(config, layer_number, enable_per_pp_rank=None):
get_recompute_priority(config, layer_number, enable_per_pp_rank=None)
num_layers_per_virtual_pipeline_stage
PipelineParallelLayerLayout
该函数由 MindSpeed 的 block recompute、norm/activation 判断、MoE zero-memory 以及 MindSpeed-LLM TransformerBlock 共同复用,避免多份公式逐渐分叉。
--swap-layer-input 不属于 CSA/MHC 重计算子图,而是可与细粒度重计算组合的独立显存优化。当前实现对组合场景做了以下增强:
layer_idx
num_layers
该增强不改变 CSA/MHC checkpoint 的输入或输出语义,也不使 CSA/MHC 重计算依赖 swap manager。其目的是确保两类显存优化在 FB overlap、多 microbatch 和 MTP 组合下独立管理各自的 Tensor 生命周期。
预期收益是减少 CSA 和 MHC 激活的峰值 NPU 显存,代价是反向前重复执行被选择的前向子图。实际收益取决于序列长度、batch size、层数、TP/CP/PP/EP 配置、CSA 压缩比例、是否启用 MHC、FP8 和 FB overlap。
本 RFC 不预设百分比指标。性能验收至少记录:
本特性只改变设备内存生命周期和计算时序,不新增网络接口、持久化文件或用户数据采集,不扩大数据访问范围。主要安全风险来自非法 storage 操作而非传统数据安全问题。
.backward()
no_grad
--use-fused-mhc
后续建议增加可选 debug 日志或 profiler 标记,记录每层启用的 checkpoint、释放时刻、重算次数和峰值显存。默认关闭日志,避免影响性能。
开发环境:Python、PyTorch/Megatron-Core、MindSpeed、MindSpeed-LLM、Ascend NPU 与 torch_npu。性能与正确性验收需在支持目标融合算子、FB overlap 和相应通信拓扑的 NPU 集群完成。
开发约束:
.grad()
可验收设计:采用固定 seed、相同数据、相同并行配置,对比 feature off、细粒度重计算和全层重计算。数值误差使用项目既有 BF16/FP16/FP32 标准;显存与性能在稳定若干 step 后采样。
recompute_csa_attention
bool
True
MEM_ARGS=" --recompute-csa-attention \ --recompute-norm \ "
mhc_recompute
--enable-mhc
MEM_ARGS=" --enable-mhc \ --use-fused-mhc \ --mhc-recompute \ --recompute-csa-attention \ "
get_recompute_priority
get_recompute_priority(config, layer_number, enable_per_pp_rank=None) -> int
config
layer_number
int
enable_per_pp_rank
Optional[bool]
None/False/True
discard_csa_attention_intermediate_outputs() -> None discard_csa_attention_output(hook_tensor: torch.Tensor) -> None
hook_tensor
torch.Tensor
当前已新增 docs/zh/pytorch/features/mcore/deepseek4_fine_grained_recompute.md,说明:
docs/zh/pytorch/features/mcore/deepseek4_fine_grained_recompute.md
特性入口已同步加入 docs/zh/pytorch/features/README.md 和 docs/zh/docs_guide.md。
docs/zh/pytorch/features/README.md
docs/zh/docs_guide.md
至少覆盖以下组合:
recompute-csa-attention
recompute-csa-attention + recompute-norm
mhc-recompute + use-fused-mhc
recompute-csa-attention + mhc-recompute
swap-layer-input
recompute-csa-attention + swap-layer-input
recompute-csa-attention + mhc-recompute + swap-layer-input
在固定 seed 和固定 batch 上分别运行:
采集并比较:
当前开发验证中已观察到细粒度重计算与全层重计算的 loss 可以对齐,但 grad norm 仍可能存在较明显差异。该项不能仅以 loss 对齐替代,应作为合入前的重点验收项,进一步区分随机数状态、算子非确定性、重算精度路径和输出释放时序的影响。
本方案提高了模型与 overlap 调度之间的协作复杂度。新增子图或改变 forward 顺序时,维护者必须重新检查“谁持有输出、何时最后一次前向读取、在哪个 Tensor 上注册 backward hook”三个问题。
recompute_method
recompute_num_layers
与通用 checkpoint 相比,本方案的主要差异是面向 DeepSeek V4 CSA 与 MHC 的模型语义划分,以及对 FB overlap、MHC post 和 o-up 特殊依赖的显式处理。
mindspeed/core/memory/recompute/recompute_common.py
mindspeed/core/memory/swap_layer_input/swap_layer_input.py
mindspeed/core/memory/swap_layer_input/swap_layer_input_manager.py
mindspeed/features_manager/memory/swap_layer_input.py
mindspeed/core/transformer/moe/moe_feature/fb_overlap/modules/utils.py
NoopLayerGraph
mindspeed/core/transformer/moe/moe_feature/fb_overlap/modules/attention.py
mindspeed/core/transformer/moe/moe_feature/fb_overlap/overlap_funcs/fwd.py
mindspeed/core/transformer/moe/moe_feature/fb_overlap/overlap_funcs/fwdbwd.py
mindspeed/core/transformer/transformer_block.py
mindspeed_llm/features_manager/transformer/multi_latent_attention/csa_feature.py
mindspeed_llm/features_manager/transformer/mhc_feature.py
mindspeed_llm/tasks/models/transformer/deepseek4/csa.py
mindspeed_llm/tasks/models/transformer/deepseek4/compressor.py
mindspeed_llm/tasks/models/transformer/deepseek4/csa_fb_overlap.py
mindspeed_llm/tasks/models/transformer/deepseek4/mhc.py
mindspeed_llm/core/transformer/transformer_layer.py
mindspeed_llm/features_manager/memory/swap_layer_input.py
mindspeed_llm/features_manager/__init__.py
欢迎加入社区,感谢您对社区的贡献 🎉!
RFC:DeepSeek V4 细粒度重计算
mhc-recompute-fboverlaprecompute-g2-attention--recompute-csa-attention、--mhc-recompute1. 概述
1.1 简介
本提案为 DeepSeek V4 Mcore 训练增加细粒度重计算能力。方案不对整个 Transformer Layer 做统一重计算,而是选择 CSA 与融合 MHC 中内存占用较高、计算代价可接受的中间结果,在前向阶段主动释放其底层存储,并在反向传播即将使用时恢复,从而降低训练峰值显存。
方案由两个仓库协同实现:Megatron/MindSpeed 提供无输出保存重计算基础能力,MindSpeed 提供 Pipeline Parallel(PP)层优先级计算以及 MoE forward-backward overlap(以下简称 FB overlap)中的释放时序;MindSpeed-LLM 提供 DeepSeek V4 模型侧的 CSA indexer、稀疏注意力、输出投影、compressor 和 MHC 重计算逻辑。用户通过
--recompute-csa-attention和--mhc-recompute两个参数控制主要能力。1.2 动机
DeepSeek V4 的 CSA 前向链路包含 q-up projection、q-norm/RoPE、indexer、压缩 KV、稀疏注意力、o-down projection 和 o-up projection;启用 MHC 后还会增加 MHC pre/post 的中间结果。在长序列、MoE、FB overlap、PP/VPP 或自定义
pipeline-model-parallel-layout组合下,这些激活可能在反向传播前长期驻留,成为峰值显存的重要组成部分。当前实现需要同时处理多个子模块之间的生命周期依赖:CSA o-up 输出会继续进入 BDA 和 MHC post;FB overlap 会拆分并交错执行前向、反向与通信;MTP 和自定义 PP layout 还会改变层构建与 layer/chunk 映射。因此,细粒度重计算不仅需要划分合适的重计算子图,还需要在普通 Transformer 与 FB overlap 路径中选择正确的输出释放点和反向挂钩点。
本方案使用
--recompute-csa-attention统一控制 CSA 子图,在内部对 q-norm、sparse attention 和 o-up 分别管理 checkpoint 生命周期,并通过 MindSpeed 的 overlap 调度协调 BDA、MHC 和 norm 重计算。对外仅注册功能级参数,不暴露子图级开关。1.3 目标
目标
--recompute-csa-attention控制 q-up、indexer、sparse attention、o-down 和 o-up 等 CSA 子图重计算。--mhc-recompute支持融合 NPU attention/MLP MHC pre 和 attention MHC post 重计算。pipeline-model-parallel-layout组合下的层选择与构建;不对 MTP 层启用当前尚未验证的 CSA/MHC 细粒度重计算。非目标
2. 用例分析
2.1 DeepSeek V4 CSA 训练
用户希望在不启用整层重计算的情况下,仅重计算 CSA 中显存占用较高的子图。开启
--recompute-csa-attention后,方案覆盖:--recompute-norm组合时的 q-norm/RoPE。关键要求是训练 loss 与不开启细粒度重计算的基线在既有数值容差内对齐,且反向梯度图完整。
2.2 CSA 与 MHC 联合重计算
当同时使用
--enable-mhc --use-fused-mhc --mhc-recompute --recompute-csa-attention时,CSA 的输出会进入 attention BDA 和 MHC post。释放点必须同时满足:(y, post, comb)与 attention MHC post 输出分别建立 checkpoint,并在安全的后续图节点上挂接恢复逻辑。2.3 MoE FB overlap
FB overlap 会拆分并交错执行不同层的前向和反向图。方案需要在 MindSpeed 的 overlap 调度中保存 checkpoint manager,并在合适的 layer output 上注册重计算 hook,同时不干扰 All-to-All、专家计算、权重梯度解耦和既有 swap 流程。
普通 MoE FB overlap 与 balanced MoE 使用不同的
attention_forward引用,因此 MindSpeed-LLM 需要通过 feature patch 将统一的 CSA 释放包装器安装到对应入口。2.4 PP/VPP、自定义 PP layout 与 MTP
细粒度重计算通常会与
recompute-num-layers、recompute-norm-num-layers等按层配置一起使用。自定义 PP layout 下,每个 PP/VPP chunk 的 decoder layer 数可能不均匀,不能再用全局平均层数推导 chunk 大小。方案要求直接解析
pipeline_model_parallel_layout,依据当前 PP rank、VPP rank 和 decoder layer 的实际位置计算重计算优先级。MTP 层单独标记并跳过当前 CSA/MHC 细粒度重计算,以避免把主模型层的生命周期假设错误应用到 MTP 图。2.5 Swap layer input 组合场景
细粒度重计算与
--swap-layer-input可以组合使用,但两者职责不同:FB overlap 会使不同层、不同 microbatch 的前反向交错。为避免仅根据相邻层和 manager 队列顺序恢复到错误输入,当前实现将每次换出产生的
swap_entry与对应LayerGraph绑定,在bwd_layer_graph或next_bwd_layer_graph进入反向前精确恢复。恢复完成后按对象身份从 manager 中移除 entry,避免扰乱其他 microbatch 的队列。MindSpeed 保留面向 Megatron TransformerLayer 的默认 feature,同时提供通用 swap wrapper 和 manager;MindSpeed-LLM 保留自己的
SwapLayerInputFeature,负责向自定义 TransformerLayer 及所需 FB overlap 入口注册 patch。MindSpeed-LLM 不重复实现 swap wrapper 和 manager。MTP Layer 不参与当前 layer-input swap manager 调度,避免重复构建影响主模型的 manager 数量和层编号。2.6 DFX 要求
False;未开启时保持原执行路径。training且torch.is_grad_enabled()时建立 checkpoint;MHC 仅在融合实现、训练态和非 MTP 层建立 checkpoint,并依赖正常训练反向图完成恢复。3. 方案设计
3.1 总体方案
整体设计分为模型策略层、通用重计算层和 overlap 生命周期层。
flowchart LR A["CLI 参数"] --> B["MindSpeed-LLM Feature Manager"] B --> C["DeepSeek V4 CSA 与 MHC 子图"] C --> D["CheckpointWithoutOutput"] D --> E["前向计算并保留 autograd 上下文"] E --> F["到达安全释放点"] F --> G["释放输出 storage 并注册 backward hook"] G --> H["反向到达 hook"] H --> I["恢复 RNG 并重算输出"] I --> J["继续原反向传播"] K["MindSpeed FB overlap 调度"] --> F L["PP layout 重计算优先级"] --> C M["Swap layer input"] --> N["LayerGraph 绑定 swap entry"] N --> O["反向前换回 Layer 输入"]职责划分如下:
--recompute-csa-attention、--mhc-recompute,决定具体子图是否重计算LayerGraph在反向前精确恢复recompute_priority3.2 技术选型
当前代码采用以下技术组件:
--recompute-csa-attention和--mhc-recompute两个功能级参数CheckpointWithoutOutput保留输入、RNG 和 autograd 上下文,并允许释放输出 storageLayerGraph的精确恢复;MindSpeed-LLM 负责自身 patch 注册pipeline_model_parallel_layout计算重计算优先级is_mtp_layer和is_mtp_attention显式跳过当前细粒度重计算3.3 功能与性能设计
3.3.1 CSA 子图划分
--recompute-csa-attention开启后,DeepSeek V4 CSA 依据运行条件建立以下 checkpoint:CheckpointWithoutOutput--recompute-norm和层选择联动CheckpointWithoutOutputkv_compress生成后释放 FP32 转换输出;FP8/TND/无压缩输出路径不启用CheckpointWithoutOutputCheckpointWithoutOutput重计算只在以下条件同时满足时启用:
3.3.2 分阶段释放策略
CSA 与 MHC post 同时重计算时,生命周期分为两个阶段:
sequenceDiagram participant CSA as CSA participant BDA as Attention BDA participant MHC as MHC Post participant L as Layer/FB Overlap participant BW as Backward CSA->>CSA: 计算 q-norm、sparse、o-up CSA->>BDA: 返回 o-up 输出 BDA->>MHC: 生成 BDA 输出并进入 MHC post MHC-->>L: 返回 MHC post 输出 L->>CSA: 立即释放 q-norm/sparse 中间输出 Note over CSA,L: o-up 暂时保留 L->>L: 继续 MLP/MHC,形成稳定 hook tensor L->>CSA: 释放 o-up,并为所有相关 checkpoint 注册 hook L->>MHC: 释放受支持的 MHC pre/post 输出并注册 hook BW->>CSA: hook 触发,按依赖顺序重算 BW->>MHC: 恢复 MHC 输出并继续反向MindSpeed-LLM 的
discard_csa_attention_intermediate_outputs()只清理 q-norm 和 sparse checkpoint 输出;discard_csa_attention_output(hook_tensor)负责最终挂接 q-norm、sparse 与 o-up 的重计算。csa_fb_overlap_attention_forward_wrapper()根据defer_attention_recompute_for_mhc_post判断采用单阶段还是分阶段释放。该设计解决两类相反问题:
3.3.3 MHC pre/post 重计算
--mhc-recompute面向融合 NPU MHC 实现,当前覆盖 attention/MLP MHC pre 以及 attention MHC post:(y, post, comb)包装为 viewless tensor,避免 view tensor 与 checkpoint/storage 操作冲突。x、residual、post、comb作为显式 checkpoint 输入;MLP MHC post 不建立当前细粒度 checkpoint。is_mtp_layer=True时不启用上述 MHC 重计算。3.3.4 BDA 与 norm 联动
当 CSA 与 MHC post 需要统一延迟挂钩时,MindSpeed 将 attention BDA 也纳入
CheckpointWithoutOutput。根据 attention bias 是否存在,分别构造无 bias 和有 bias 的 checkpoint 函数,保持 BDA 调用签名稳定。--recompute-norm开启时:defer_attention_recompute_norm,由统一释放函数在安全 hook tensor 上注册 norm 重计算。3.3.5 FB overlap 适配
MindSpeed 的 FB overlap forward/backward 图增加以下能力:
LayerGraph中保留既有 activation checkpoint manager 与 swap manager,不改变专家通信和权重梯度解耦顺序。MindSpeed-LLM 仅在以下条件成立时注册 FB overlap patch:
balanced MoE 开启时额外 patch balanced MoE 的 forward 与 forward-backward 入口。包装器带
_csa_fb_overlap_wrapped标识,防止重复包装。3.3.6 PP layout 重计算优先级
MindSpeed 提供统一的
get_recompute_priority(config, layer_number, enable_per_pp_rank=None):num_layers_per_virtual_pipeline_stage为空时提供基础回退计算。pipeline_model_parallel_layout时,将字符串或列表解析为PipelineParallelLayerLayout。该函数由 MindSpeed 的 block recompute、norm/activation 判断、MoE zero-memory 以及 MindSpeed-LLM TransformerBlock 共同复用,避免多份公式逐渐分叉。
3.3.7 Swap layer input 兼容增强
--swap-layer-input不属于 CSA/MHC 重计算子图,而是可与细粒度重计算组合的独立显存优化。当前实现对组合场景做了以下增强:SwapLayerInputFeature只 patch Megatron TransformerLayer;MindSpeed-LLM 保留本地SwapLayerInputFeature,导入 MindSpeed wrapper 并 patch 自身 TransformerLayer 和 FB overlap 入口。swap_entry,并将其保存在本次前向生成的LayerGraph。bwd_layer_graph和next_bwd_layer_graph对应的精确 entry,不再只依赖相邻层和 FIFO 位置推断 microbatch。layer_idx和num_layers与主模型不一致。该增强不改变 CSA/MHC checkpoint 的输入或输出语义,也不使 CSA/MHC 重计算依赖 swap manager。其目的是确保两类显存优化在 FB overlap、多 microbatch 和 MTP 组合下独立管理各自的 Tensor 生命周期。
3.3.8 性能影响
预期收益是减少 CSA 和 MHC 激活的峰值 NPU 显存,代价是反向前重复执行被选择的前向子图。实际收益取决于序列长度、batch size、层数、TP/CP/PP/EP 配置、CSA 压缩比例、是否启用 MHC、FP8 和 FB overlap。
本 RFC 不预设百分比指标。性能验收至少记录:
3.4 安全隐私与 DFX 设计
安全与隐私
本特性只改变设备内存生命周期和计算时序,不新增网络接口、持久化文件或用户数据采集,不扩大数据访问范围。主要安全风险来自非法 storage 操作而非传统数据安全问题。
可靠性
.backward()训练流程,不把训练态下的no_grad执行作为支持场景。兼容性
--mhc-recompute只有在--use-fused-mhc的训练路径上有效;非融合 MHC 保持原逻辑。--swap-layer-input与 CSA/MHC 重计算分别管理 Layer 输入和子图内部激活;两者可以同时开启。LayerGraph绑定,避免跨层或跨 microbatch 错配。可维护性
可观测性建议
后续建议增加可选 debug 日志或 profiler 标记,记录每层启用的 checkpoint、释放时刻、重算次数和峰值显存。默认关闭日志,避免影响性能。
3.5 编程与调用设计
3.5.1 编程模型基本设计
开发环境:Python、PyTorch/Megatron-Core、MindSpeed、MindSpeed-LLM、Ascend NPU 与 torch_npu。性能与正确性验收需在支持目标融合算子、FB overlap 和相应通信拓扑的 NPU 集群完成。
开发约束:
CheckpointWithoutOutput不支持.grad()风格的 checkpoint 反向,使用训练框架标准.backward()。可验收设计:采用固定 seed、相同数据、相同并行配置,对比 feature off、细粒度重计算和全层重计算。数值误差使用项目既有 BF16/FP16/FP32 标准;显存与性能在稳定若干 step 后采样。
3.5.2 接口定义与设计
3.5.2.1
--recompute-csa-attention--recompute-csa-attentionrecompute_csa_attentionboolFalse(默认)/True--recompute-norm且当前层被 norm 重计算策略选中时纳入 CSA checkpoint。MEM_ARGS=" --recompute-csa-attention \ --recompute-norm \ "3.5.2.2
--mhc-recompute--mhc-recomputemhc_recomputeboolFalse(默认)/True--enable-mhc与--use-fused-mhc;MTP 层跳过,MLP MHC post 不在当前重计算范围内。MEM_ARGS=" --enable-mhc \ --use-fused-mhc \ --mhc-recompute \ --recompute-csa-attention \ "3.5.2.3
get_recompute_priorityget_recompute_priority(config, layer_number, enable_per_pp_rank=None) -> intconfiglayer_numberintenable_per_pp_rankOptional[bool]None/False/Truerecompute_priorityint3.5.2.4 CSA 输出生命周期接口
discard_csa_attention_intermediate_outputs() -> None discard_csa_attention_output(hook_tensor: torch.Tensor) -> Nonehook_tensortorch.Tensor3.5.3 用户文档
当前已新增
docs/zh/pytorch/features/mcore/deepseek4_fine_grained_recompute.md,说明:--recompute-csa-attention与--mhc-recompute的作用和默认值。--enable-mhc、--use-fused-mhc的依赖。特性入口已同步加入
docs/zh/pytorch/features/README.md和docs/zh/docs_guide.md。4. 测试设计
4.1 单元测试
bwd_layer_graph/next_bwd_layer_graph恢复正确4.2 集成测试
至少覆盖以下组合:
recompute-csa-attention单独开启。recompute-csa-attention + recompute-norm。mhc-recompute + use-fused-mhc。recompute-csa-attention + mhc-recompute。pipeline-model-parallel-layout。swap-layer-input单独开启,验证基线换入换出行为。recompute-csa-attention + swap-layer-input,并叠加多 microbatch 和 MoE FB overlap。recompute-csa-attention + mhc-recompute + swap-layer-input,并叠加 MTP、PP/VPP 和非均匀 PP layout。4.3 端到端测试
在固定 seed 和固定 batch 上分别运行:
采集并比较:
当前开发验证中已观察到细粒度重计算与全层重计算的 loss 可以对齐,但 grad norm 仍可能存在较明显差异。该项不能仅以 loss 对齐替代,应作为合入前的重点验收项,进一步区分随机数状态、算子非确定性、重算精度路径和输出释放时序的影响。
4.4 回归与静态检查
5. 缺点和风险
LayerGraph显式绑定,同时覆盖 current/next backward graphLayerGraph引用is_mtp_layer/is_mtp_attention显式跳过本方案提高了模型与 overlap 调度之间的协作复杂度。新增子图或改变 forward 顺序时,维护者必须重新检查“谁持有输出、何时最后一次前向读取、在哪个 Tensor 上注册 backward hook”三个问题。
6. 现有技术
CheckpointWithoutOutput:允许前向保留 Tensor 对象与 autograd 上下文,同时清空输出 storage,并在反向前重算回填。本方案将其用于 CSA、MHC、BDA 和 compressor,并补充跨 FB overlap 的生命周期管理。recompute_method、recompute_num_layers等选择整层。本方案与其并存,并统一复用 PP layout-aware 的层优先级。LayerGraph绑定独立 swap entry,保证 FB overlap 和多 microbatch 下的精确恢复。MindSpeed-LLM 保留模型侧 patch 注册,底层换入换出能力复用 MindSpeed。与通用 checkpoint 相比,本方案的主要差异是面向 DeepSeek V4 CSA 与 MHC 的模型语义划分,以及对 FB overlap、MHC post 和 o-up 特殊依赖的显式处理。
7. 未解决问题
附录
A. 主要代码位置
MindSpeed
mindspeed/core/memory/recompute/recompute_common.py:无输出保存重计算与 PP layout priority。mindspeed/core/memory/swap_layer_input/swap_layer_input.py:TransformerLayer wrapper、FB overlap 换出与基于LayerGraph的精确恢复。mindspeed/core/memory/swap_layer_input/swap_layer_input_manager.py:独立 swap entry 追踪、指定 entry 等待/恢复和按身份移除。mindspeed/features_manager/memory/swap_layer_input.py:MindSpeed 侧SwapLayerInputFeature,负责 Megatron TransformerLayer 的默认 patch 注册。mindspeed/core/transformer/moe/moe_feature/fb_overlap/modules/utils.py:在LayerGraph/NoopLayerGraph中保存对应 swap entry。mindspeed/core/transformer/moe/moe_feature/fb_overlap/modules/attention.py:CSA、norm、BDA、MHC 的统一释放时序。mindspeed/core/transformer/moe/moe_feature/fb_overlap/overlap_funcs/fwd.py:FB overlap 前向集成。mindspeed/core/transformer/moe/moe_feature/fb_overlap/overlap_funcs/fwdbwd.py:交错前反向集成。mindspeed/core/transformer/transformer_block.py:block recompute 与 layout-aware priority。MindSpeed-LLM
mindspeed_llm/features_manager/transformer/multi_latent_attention/csa_feature.py:CSA 参数及 FB overlap patch 注册。mindspeed_llm/features_manager/transformer/mhc_feature.py:MHC 重计算参数。mindspeed_llm/tasks/models/transformer/deepseek4/csa.py:CSA 子图重计算与分阶段释放。mindspeed_llm/tasks/models/transformer/deepseek4/compressor.py:compressor FP32 转换重计算。mindspeed_llm/tasks/models/transformer/deepseek4/csa_fb_overlap.py:CSA FB overlap 包装器。mindspeed_llm/tasks/models/transformer/deepseek4/mhc.py:融合 attention/MLP MHC pre 与 attention MHC post 重计算。mindspeed_llm/core/transformer/transformer_layer.py:普通 Transformer 路径的释放挂钩点与 MTP 标识传递。mindspeed_llm/features_manager/memory/swap_layer_input.py:复用 MindSpeed wrapper,注册 MindSpeed-LLM TransformerLayer 与 FB overlap patch。mindspeed_llm/features_manager/__init__.py:将 MindSpeed-LLM 本地SwapLayerInputFeature加入 swap feature 列表。B. 术语表
C. 文档状态与后续补充
docs/zh/pytorch/features/mcore/deepseek4_fine_grained_recompute.md,并加入特性列表与文档目录。欢迎加入社区,感谢您对社区的贡献 🎉!