已关闭
[RFC]: DSK V4细粒度重计算 #289
wuweiqiang24创建于  8月13日关闭于  3 天前
wuweiqiang24
wuweiqiang24成员
8月13日 创建

RFC:DeepSeek V4 细粒度重计算

项目 内容
涉及仓库 MindSpeed、MindSpeed-LLM
MindSpeed 开发分支 mhc-recompute-fboverlap
MindSpeed-LLM 开发分支 recompute-g2-attention
目标模型 DeepSeek V4
目标硬件 Ascend NPU
对外参数 --recompute-csa-attention--mhc-recompute
文档状态 Draft(按当前实现更新)

MindSpeed-LLM 分支名沿用开发初期名称,其中的 g2-attention 仅是历史分支标识。当前代码、接口和文档统一使用 CSA 命名。

1. 概述

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 子图重计算。
  • 对 CSA 的高内存中间结果实施选择性重计算,保持模型输入、输出和参数语义不变。
  • 通过 --mhc-recompute 支持融合 NPU attention/MLP MHC pre 和 attention MHC post 重计算。
  • 支持普通 Transformer forward、MoE FB overlap 以及 balanced MoE FB overlap 路径。
  • 在 CSA 与 MHC post 同时开启时,保证释放时序正确:q-norm/sparse 等中间结果尽早释放,o-up 仅延迟到安全挂钩点。
  • 支持 PP、VPP、MTP 和 pipeline-model-parallel-layout 组合下的层选择与构建;不对 MTP 层启用当前尚未验证的 CSA/MHC 细粒度重计算。
  • 默认关闭特性,未开启相关参数时不改变现有训练路径。
  • 对外只保留功能级参数,不继续暴露 q-up、indexer、sparse attention、o-down 和 o-up 的独立开关。

非目标

  • 不替代 Megatron/MindSpeed 已有的全层重计算、uniform/block 重计算或其他模型的 MLA 重计算。
  • 不保证在所有模型、所有 Attention 实现或非 Ascend 设备上自动生效。
  • 不在本提案中修改模型数学公式、优化器、通信算法或训练数据格式。
  • 不在本提案中承诺固定的显存下降比例或吞吐损失;量化指标需要在目标模型配置和 NPU 集群上实测。
  • 当前不对 MTP 层启用 CSA/MHC 细粒度重计算。

2. 用例分析

2.1 DeepSeek V4 CSA 训练

用户希望在不启用整层重计算的情况下,仅重计算 CSA 中显存占用较高的子图。开启 --recompute-csa-attention 后,方案覆盖:

  • q-up projection,以及与 --recompute-norm 组合时的 q-norm/RoPE。
  • CSA indexer 的索引与分数计算。
  • compressor 中非 FP8 路径的 FP32 输入转换结果。
  • 带或不带 fused lightning indexer loss 的 sparse attention。
  • o-down、输出 RoPE、reshape 和 o-up projection。

关键要求是训练 loss 与不开启细粒度重计算的基线在既有数值容差内对齐,且反向梯度图完整。

2.2 CSA 与 MHC 联合重计算

当同时使用 --enable-mhc --use-fused-mhc --mhc-recompute --recompute-csa-attention 时,CSA 的输出会进入 attention BDA 和 MHC post。释放点必须同时满足:

  • 已建立 checkpoint 的 q-norm、sparse attention 等不再被后续前向算子直接读取的中间结果应尽快释放。
  • o-up 输出不能与 q-norm/sparse 一起盲目清空,应保留到 MHC post 输出已经形成、能够注册反向重计算 hook 的安全时刻。
  • attention/MLP MHC pre 的 (y, post, comb) 与 attention MHC post 输出分别建立 checkpoint,并在安全的后续图节点上挂接恢复逻辑。
  • MHC 重计算仅在模块训练态、使用融合 NPU MHC 且非 MTP 层时生效;实际反向仍要求处于正常训练计算图中。

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-layersrecompute-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 可以组合使用,但两者职责不同:

  • CSA/MHC 细粒度重计算释放模型子图的中间激活,在反向前重新计算。
  • swap layer input 将 Transformer Layer 输入异步换出到 CPU,并在对应反向图运行前换回 NPU。
  • 两者管理不同 Tensor,CSA/MHC checkpoint 不依赖 swap manager,但在 FB overlap 与多 microbatch 场景中需要协调生命周期。

FB overlap 会使不同层、不同 microbatch 的前反向交错。为避免仅根据相邻层和 manager 队列顺序恢复到错误输入,当前实现将每次换出产生的 swap_entry 与对应 LayerGraph 绑定,在 bwd_layer_graphnext_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;未开启时保持原执行路径。
  • 可维护性:用户配置以两个语义化参数为主,内部 checkpoint manager 按子图职责命名。
  • 可靠性:CSA 仅在 trainingtorch.is_grad_enabled() 时建立 checkpoint;MHC 仅在融合实现、训练态和非 MTP 层建立 checkpoint,并依赖正常训练反向图完成恢复。
  • 可测试性:分别验证单特性、组合特性、FB overlap、MTP、PP layout 和数值一致性。
  • 可诊断性:对非法 PP layout、缺失 layer number、超出 VPP rank 范围等情况抛出明确异常。

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 输入"]

职责划分如下:

层次 仓库 职责
参数与模型策略 MindSpeed-LLM 注册 --recompute-csa-attention--mhc-recompute,决定具体子图是否重计算
模型子图 MindSpeed-LLM 划分 q-up/q-norm、indexer、compressor、sparse attention、o-down/o-up、MHC pre 及 attention MHC post
swap patch 注册 MindSpeed-LLM 为自定义 TransformerLayer 和所需 FB overlap 入口注册 swap wrapper
通用 checkpoint Megatron/MindSpeed 保存输入与 RNG 状态,释放输出 storage,反向前重算并回填输出
FB overlap 生命周期 MindSpeed 在拆分图中选择安全 hook tensor,协调 CSA、norm、BDA、MHC 和 swap manager
Layer 输入换入换出 MindSpeed 为每次换出追踪独立 entry,并按 LayerGraph 在反向前精确恢复
层选择 MindSpeed 根据普通 PP/VPP 或实际 PP layout 计算 recompute_priority

3.2 技术选型

当前代码采用以下技术组件:

技术组件 当前实现
控制接口 使用 --recompute-csa-attention--mhc-recompute 两个功能级参数
checkpoint 机制 使用 CheckpointWithoutOutput 保留输入、RNG 和 autograd 上下文,并允许释放输出 storage
CSA 生命周期 q-norm/sparse 中间输出与 o-up 输出分阶段释放
MHC 生命周期 attention/MLP MHC pre 与 attention MHC post 分别建立 checkpoint,并在 layer 或 FB overlap 的安全 Tensor 上注册 hook
overlap 集成 MindSpeed 负责 FB overlap 中 CSA、norm、BDA、MHC 和 swap manager 的协调
swap input 集成 MindSpeed 提供 wrapper、manager、独立 swap entry 和基于 LayerGraph 的精确恢复;MindSpeed-LLM 负责自身 patch 注册
模型集成 MindSpeed-LLM 负责 DeepSeek V4 CSA、compressor 和 MHC 子图划分
PP 层选择 根据普通 PP/VPP 配置或实际 pipeline_model_parallel_layout 计算重计算优先级
MTP 处理 通过 is_mtp_layeris_mtp_attention 显式跳过当前细粒度重计算

3.3 功能与性能设计

3.3.1 CSA 子图划分

--recompute-csa-attention 开启后,DeepSeek V4 CSA 依据运行条件建立以下 checkpoint:

子图 实现方式 释放/恢复时机
q-up projection CheckpointWithoutOutput 未命中 q-norm 重计算时,在 q-norm/RoPE 已生成后释放 q-up 输出,反向到 q 时恢复
q-up + q-norm + RoPE --recompute-norm 和层选择联动 命中 q-norm 重计算时合并为一个 checkpoint;无 sparse checkpoint 时直接挂到 attention 输出,有 sparse checkpoint 时交由统一 CSA 生命周期管理
indexer 标准 tensor-parallel checkpoint 反向阶段重新执行 indexer 子图
compressor FP32 转换 CheckpointWithoutOutput kv_compress 生成后释放 FP32 转换输出;FP8/TND/无压缩输出路径不启用
sparse attention CheckpointWithoutOutput q-norm/sparse 阶段尽早释放,反向前恢复
o-down + 输出 RoPE + o-up CheckpointWithoutOutput 普通路径在 BDA 输出处释放;CSA+MHC+FB overlap 时延迟到 MHC post 安全点

重计算只在以下条件同时满足时启用:

参数开启 AND 非 MTP Attention AND module.training AND torch.is_grad_enabled()

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 判断采用单阶段还是分阶段释放。

该设计解决两类相反问题:

  • 不清除或过晚清除 q-norm/sparse 会扩大前向峰值显存。
  • 过早清除 o-up 会破坏后续 BDA/MHC 所需的输出或反向依赖。

3.3.3 MHC pre/post 重计算

--mhc-recompute 面向融合 NPU MHC 实现,当前覆盖 attention/MLP MHC pre 以及 attention MHC post:

  • MHC pre 将融合算子输出 (y, post, comb) 包装为 viewless tensor,避免 view tensor 与 checkpoint/storage 操作冲突。
  • attention MHC post 将 x、residual、post、comb 作为显式 checkpoint 输入;MLP MHC post 不建立当前细粒度 checkpoint。
  • 普通 Transformer 路径在 layer 最终输出上释放 attention/MLP MHC pre 与 attention MHC post。
  • FB overlap 路径由 MindSpeed 在 MLP MHC post 输出或 dense layer 输出上统一释放。
  • is_mtp_layer=True 时不启用上述 MHC 重计算。
  • FB overlap 且未开启 CSA 重计算时,MHC post 对瞬时 attention 输出执行必要的 clone,避免该输入在 checkpoint hook 运行前被 overlap 路径释放。

3.3.4 BDA 与 norm 联动

当 CSA 与 MHC post 需要统一延迟挂钩时,MindSpeed 将 attention BDA 也纳入 CheckpointWithoutOutput。根据 attention bias 是否存在,分别构造无 bias 和有 bias 的 checkpoint 函数,保持 BDA 调用签名稳定。

--recompute-norm 开启时:

  • 普通场景仍在 BDA 输出上释放 norm checkpoint。
  • 需要等待 MHC post 的场景记录 defer_attention_recompute_norm,由统一释放函数在安全 hook tensor 上注册 norm 重计算。
  • CSA q-norm 与 MindSpeed 通用 norm 重计算通过同一 layer priority 逻辑选择层,避免 PP layout 下策略不一致。

3.3.5 FB overlap 适配

MindSpeed 的 FB overlap forward/backward 图增加以下能力:

  • 在 attention、MLP MHC post 和 layer output 处调用统一释放函数。
  • 将 CSA、受支持的 MHC pre/post、norm 和 BDA checkpoint 的 hook 绑定到真实存活的图节点。
  • LayerGraph 中保留既有 activation checkpoint manager 与 swap manager,不改变专家通信和权重梯度解耦顺序。
  • 同时覆盖普通 FB overlap 和 balanced MoE overlap 的 attention forward 引用。

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 标识,防止重复包装。

3.3.6 PP layout 重计算优先级

MindSpeed 提供统一的 get_recompute_priority(config, layer_number, enable_per_pp_rank=None)

  • 无自定义 layout 时,兼容普通 PP/VPP,并在 num_layers_per_virtual_pipeline_stage 为空时提供基础回退计算。
  • 存在 pipeline_model_parallel_layout 时,将字符串或列表解析为 PipelineParallelLayerLayout
  • 按 layout 中 decoder layer 的实际分布,构造当前 PP rank 在各 VPP chunk 中的 layer id。
  • 未开启 per-PP-rank 策略时返回当前 chunk 内局部位置;开启后按各 chunk 相同局部深度交错计算 priority。
  • 对非法 layout 类型、layer 不属于当前 rank/chunk、VPP rank 越界等情况抛出异常。

该函数由 MindSpeed 的 block recompute、norm/activation 判断、MoE zero-memory 以及 MindSpeed-LLM TransformerBlock 共同复用,避免多份公式逐渐分叉。

3.3.7 Swap layer input 兼容增强

--swap-layer-input 不属于 CSA/MHC 重计算子图,而是可与细粒度重计算组合的独立显存优化。当前实现对组合场景做了以下增强:

  • MindSpeed 的 SwapLayerInputFeature 只 patch Megatron TransformerLayer;MindSpeed-LLM 保留本地 SwapLayerInputFeature,导入 MindSpeed wrapper 并 patch 自身 TransformerLayer 和 FB overlap 入口。
  • 每次有效的 D2H 换出返回独立 swap_entry,并将其保存在本次前向生成的 LayerGraph
  • FB overlap 的交错调度在执行反向前恢复 bwd_layer_graphnext_bwd_layer_graph 对应的精确 entry,不再只依赖相邻层和 FIFO 位置推断 microbatch。
  • entry 恢复完成后按对象身份从 manager 队列中移除,避免重复恢复、队列错位和 CPU/NPU 内存泄漏。
  • MTP Layer 显式跳过 layer-input swap,避免 MTP 重复构建 TransformerLayer 导致 manager 数量、layer_idxnum_layers 与主模型不一致。
  • 对无 manager、非 Tensor 输入、无有效换出 Tensor 和末层不换出等情况做空路径保护。

该增强不改变 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 不预设百分比指标。性能验收至少记录:

  • 特性关闭、全层重计算、细粒度重计算三组峰值 NPU 显存。
  • 单 step 时间、吞吐和重计算额外算力开销。
  • q-norm/sparse 立即释放前后的峰值差异。
  • CSA 与 MHC 联合开启时 o-up 延迟释放对正确性和峰值的影响。
  • 开启 swap layer input 后的 D2H/H2D 传输量、CPU pinned memory、同步等待时间和整体吞吐。

3.4 安全隐私与 DFX 设计

安全与隐私

本特性只改变设备内存生命周期和计算时序,不新增网络接口、持久化文件或用户数据采集,不扩大数据访问范围。主要安全风险来自非法 storage 操作而非传统数据安全问题。

可靠性

  • CSA 分阶段释放前检查 checkpoint 和输出是否存在;MHC/FB overlap 路径在可用的后续图节点上注册恢复逻辑。
  • 当前功能面向标准 .backward() 训练流程,不把训练态下的 no_grad 执行作为支持场景。
  • checkpoint 保存并恢复 CPU RNG、NPU RNG 和 tensor-parallel RNG tracker 状态,避免 dropout 等随机算子因重算改变结果。
  • 重算后恢复原输出 storage 并复制结果,使下游 autograd 节点继续引用原 Tensor 对象。
  • 完成一次 backward 后清理 checkpoint manager 引用,避免跨 microbatch 复用旧上下文。

兼容性

  • 两个参数均默认关闭,对外不提供子图级兼容别名。
  • --mhc-recompute 只有在 --use-fused-mhc 的训练路径上有效;非融合 MHC 保持原逻辑。
  • MTP 层显式跳过 CSA/MHC 细粒度重计算。
  • compressor FP8 路径不做 FP32 conversion checkpoint。
  • 非 DeepSeek V4 模型不会注册 CSA FB overlap patch。
  • --swap-layer-input 与 CSA/MHC 重计算分别管理 Layer 输入和子图内部激活;两者可以同时开启。
  • FB overlap 场景中,swap entry 与实际 LayerGraph 绑定,避免跨层或跨 microbatch 错配。

可维护性

  • 通用能力放在 MindSpeed,模型语义放在 MindSpeed-LLM。
  • MindSpeed-LLM 维护自身 swap patch 注册,但直接复用 MindSpeed wrapper 和 manager,不维护第二套核心换入换出实现。
  • 对外参数按功能域聚合,内部仍保留独立 checkpoint manager,便于定位是哪段子图发生问题。
  • PP layout 层选择集中到一个公共函数,消除多个模块中的重复公式。

可观测性建议

后续建议增加可选 debug 日志或 profiler 标记,记录每层启用的 checkpoint、释放时刻、重算次数和峰值显存。默认关闭日志,避免影响性能。

3.5 编程与调用设计

3.5.1 编程模型基本设计

开发环境:Python、PyTorch/Megatron-Core、MindSpeed、MindSpeed-LLM、Ascend NPU 与 torch_npu。性能与正确性验收需在支持目标融合算子、FB overlap 和相应通信拓扑的 NPU 集群完成。

开发约束

  • MindSpeed、Megatron 与 MindSpeed-LLM 分支必须配套。
  • CheckpointWithoutOutput 不支持 .grad() 风格的 checkpoint 反向,使用训练框架标准 .backward()
  • 被释放的输出在 hook 触发前不能被未建模的异步消费者继续读取。
  • 新增重计算子图时必须将所有影响重算结果的 Tensor 作为显式输入,避免依赖已释放的闭包变量。
  • 涉及随机算子时必须纳入 RNG 保存和恢复。

可验收设计:采用固定 seed、相同数据、相同并行配置,对比 feature off、细粒度重计算和全层重计算。数值误差使用项目既有 BF16/FP16/FP32 标准;显存与性能在稳定若干 step 后采样。

3.5.2 接口定义与设计

3.5.2.1 --recompute-csa-attention
  • 接口描述:启用 DeepSeek V4 CSA 的统一细粒度重计算。
  • 接口原型--recompute-csa-attention
参数名称 输入/输出 类型 描述 取值范围
recompute_csa_attention 输入 bool 是否启用 CSA 子图重计算 False(默认)/True
  • 返回参数:无;模型输出签名保持不变。
  • 异常处理:PP layout 非法、layer 映射错误或 checkpoint 使用方式非法时抛出明确异常。
  • 约束说明:仅对受支持的 DeepSeek V4 CSA 生效;MTP Attention 跳过。q-norm 仅在同时开启 --recompute-norm 且当前层被 norm 重计算策略选中时纳入 CSA checkpoint。
  • 变更说明:该参数统一控制当前实现支持的 CSA 重计算子图。
  • 调用参考代码
MEM_ARGS="
    --recompute-csa-attention \
    --recompute-norm \
"
3.5.2.2 --mhc-recompute
  • 接口描述:启用融合 NPU MHC 的细粒度重计算,当前覆盖 attention/MLP MHC pre 与 attention MHC post。
  • 接口原型--mhc-recompute
参数名称 输入/输出 类型 描述 取值范围
mhc_recompute 输入 bool 是否重计算受支持的融合 MHC pre/post 输出 False(默认)/True
  • 返回参数:无;MHC 输出签名保持不变。
  • 异常处理:无新增用户态异常;内部仅在 checkpoint 生命周期不完整时暴露框架错误。
  • 约束说明:需要同时开启 --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_priority
  • 接口描述:为普通 PP/VPP 或自定义 PP layout 返回当前 decoder layer 的重计算优先级。
  • 接口原型
get_recompute_priority(config, layer_number, enable_per_pp_rank=None) -> int
参数名称 输入/输出 类型 描述 取值范围
config 输入 配置对象 PP/VPP、layout 和层数配置 合法 Megatron/MindSpeed 配置
layer_number 输入 int 一基 decoder layer 编号 非空正整数
enable_per_pp_rank 输入 Optional[bool] 是否按 PP rank 内跨 VPP chunk 排优先级 None/False/True
返回参数 类型 描述 取值范围
recompute_priority int 当前层在重计算选择策略中的优先级 非负整数
  • 异常处理:layer number 为空、layout 类型非法、VPP rank 越界或 layer 不属于当前 layout 时抛出异常。
  • 约束说明:layout 中的 decoder layer 分布必须与模型实际构建一致。
3.5.2.4 CSA 输出生命周期接口
  • 接口描述:供普通 Transformer 和 FB overlap 调度控制 CSA checkpoint 输出的分阶段释放。
  • 接口原型
discard_csa_attention_intermediate_outputs() -> None
discard_csa_attention_output(hook_tensor: torch.Tensor) -> None
参数名称 输入/输出 类型 描述 取值范围
hook_tensor 输入 torch.Tensor 反向到达时触发输出恢复的安全图节点 需要梯度的 Tensor
  • 约束说明:第一个接口只释放 q-norm/sparse 中间结果;第二个接口完成 q-norm、sparse、o-up 的最终 hook 注册。

3.5.3 用户文档

当前已新增 docs/zh/pytorch/features/mcore/deepseek4_fine_grained_recompute.md,说明:

  1. --recompute-csa-attention--mhc-recompute 的作用和默认值。
  2. 两个参数单独开启及组合开启的配置示例。
  3. MHC 重计算对 --enable-mhc--use-fused-mhc 的依赖。
  4. 训练态、MTP 和额外计算开销等限制。

特性入口已同步加入 docs/zh/pytorch/features/README.mddocs/zh/docs_guide.md

4. 测试设计

4.1 单元测试

测试项 验证内容
参数注册 参数默认关闭、命令行可解析、开启后正确传递到模型与 patch 注册逻辑
CSA 生效条件 train/eval、grad enabled/disabled、主层/MTP 层组合
q-up/q-norm 输出与基线一致,释放后反向可恢复
indexer fused/non-fused indexer loss、带/不带 attention mask
sparse attention 有/无 kv compress、有/无 topk idx 两条路径
compressor FP8、非 FP8、TND、无压缩输出分支
o-down/o-up bias 处理、RoPE、reshape 和反向梯度
MHC attention/MLP pre、attention post、fused/non-fused、MTP 排除及 MLP post 不重计算
生命周期 intermediate 释放幂等、最终释放清空 manager、无梯度 Tensor 不注册 hook
PP layout 均匀/非均匀 layout、PP/VPP rank、非法 layer 映射
swap entry 生命周期 D2H/H2D 状态转换、按身份移除、重复恢复幂等、空 entry 保护
swap 与 LayerGraph 绑定 不同层和 microbatch 的 entry 独立,bwd_layer_graph/next_bwd_layer_graph 恢复正确

4.2 集成测试

至少覆盖以下组合:

  1. recompute-csa-attention 单独开启。
  2. recompute-csa-attention + recompute-norm
  3. mhc-recompute + use-fused-mhc
  4. recompute-csa-attention + mhc-recompute
  5. 上述组合叠加 MoE FB overlap。
  6. balanced MoE FB overlap。
  7. PP/VPP 与非均匀 pipeline-model-parallel-layout
  8. MTP 开启,确认主模型层生效而 MTP 层跳过。
  9. CP 的 kv-allgather、TND/reset-attention-mask 路径。
  10. FP8 与非 FP8 路径。
  11. swap-layer-input 单独开启,验证基线换入换出行为。
  12. recompute-csa-attention + swap-layer-input,并叠加多 microbatch 和 MoE FB overlap。
  13. recompute-csa-attention + mhc-recompute + swap-layer-input,并叠加 MTP、PP/VPP 和非均匀 PP layout。

4.3 端到端测试

在固定 seed 和固定 batch 上分别运行:

  • 无重计算基线。
  • 细粒度 CSA/MHC 重计算。
  • 全层重计算参考组。

采集并比较:

  • 前若干 step 的 loss。
  • 每个 step 的 grad norm。
  • optimizer update 后关键参数或参数 checksum。
  • 峰值 NPU 显存、reserved/allocated memory。
  • step time、samples/s 或 tokens/s。
  • 是否出现 NaN/Inf、storage size 为 0、重复 hook 或跨 microbatch 状态泄漏。
  • swap entry 是否与对应 LayerGraph 一一匹配,反向后 manager 队列、CPU buffer 与 NPU storage 是否按预期释放。

当前开发验证中已观察到细粒度重计算与全层重计算的 loss 可以对齐,但 grad norm 仍可能存在较明显差异。该项不能仅以 loss 对齐替代,应作为合入前的重点验收项,进一步区分随机数状态、算子非确定性、重算精度路径和输出释放时序的影响。

4.4 回归与静态检查

  • 两个仓库均执行目标分支变更文件的 Ruff、Pylint 和 pre-commit;当前 Windows 工作环境缺少对应命令时,在配套 NPU/CI 环境补跑。
  • 执行现有 Transformer、MoE FB overlap、MHC、CSA/DSA、PP layout 与 MTP 测试。
  • 未开启相关参数时执行一组基线回归,确认代码路径与性能无变化。

5. 缺点和风险

风险 影响 应对措施
重算增加计算量 step time 上升 只选择高显存、可接受重算成本的子图;实测吞吐后决定默认推荐组合
输出释放过早 backward 访问空 storage、运行时报错 明确安全 hook tensor;o-up 与 q-norm/sparse 分阶段释放
输出释放过晚 峰值显存回升甚至 OOM intermediate 阶段立即释放 q-norm/sparse;profiler 验证驻留区间
RNG 状态不一致 loss/grad norm 不对齐 保存并恢复 CPU/NPU/TP RNG;增加 dropout 场景测试
算子非确定性或精度路径差异 grad norm 偏差 对比重算前后 dtype、融合/非融合算子和确定性配置
FB overlap 图复杂 hook 次序错误、状态跨层串扰 checkpoint manager 挂在 layer/module 实例;每次使用后置空;覆盖多 microbatch 测试
swap entry 与反向图错配 恢复其他 microbatch 的 Layer 输入,导致数值错误或空 storage entry 与 LayerGraph 显式绑定,同时覆盖 current/next backward graph
swap entry 重复恢复或遗留 CPU/NPU 内存泄漏、manager 队列错位 恢复后按对象身份移除 entry 并清空 LayerGraph 引用
swap 生效范围扩大 D2H/H2D 通信和同步开销增加,吞吐下降 分别测量 swap 单开和与重计算组合时的传输量、等待时间与峰值显存
MTP 生命周期不同 错误复用主层逻辑 is_mtp_layer/is_mtp_attention 显式跳过
PP layout 映射错误 选错重计算层或初始化失败 从实际 layout 获取 decoder 分布并校验 layer 所属 rank/chunk
跨仓版本耦合 单仓升级后接口缺失 明确配套分支/版本,在 patch 前使用能力检测,并做联合 CI

本方案提高了模型与 overlap 调度之间的协作复杂度。新增子图或改变 forward 顺序时,维护者必须重新检查“谁持有输出、何时最后一次前向读取、在哪个 Tensor 上注册 backward hook”三个问题。

6. 现有技术

  • Megatron activation checkpointing:支持整段函数或整层重计算,提供基础 RNG 和 autograd 机制。本方案复用其思想,但把范围缩小到 DeepSeek V4 特定子图。
  • MindSpeed CheckpointWithoutOutput:允许前向保留 Tensor 对象与 autograd 上下文,同时清空输出 storage,并在反向前重算回填。本方案将其用于 CSA、MHC、BDA 和 compressor,并补充跨 FB overlap 的生命周期管理。
  • MindSpeed 全层/按层重计算:通过 recompute_methodrecompute_num_layers 等选择整层。本方案与其并存,并统一复用 PP layout-aware 的层优先级。
  • Swap activation/layer input:通过 NPU 与 CPU 间异步换入换出节省显存,消耗通信带宽与 CPU pinned memory。细粒度重计算以额外计算换显存,两者可组合使用但职责不同。当前实现进一步以 LayerGraph 绑定独立 swap entry,保证 FB overlap 和多 microbatch 下的精确恢复。MindSpeed-LLM 保留模型侧 patch 注册,底层换入换出能力复用 MindSpeed。

与通用 checkpoint 相比,本方案的主要差异是面向 DeepSeek V4 CSA 与 MHC 的模型语义划分,以及对 FB overlap、MHC post 和 o-up 特殊依赖的显式处理。

7. 未解决问题

  1. 需要确定目标 DeepSeek V4 配置上的显存下降和吞吐损失验收阈值。
  2. 细粒度重计算与全层重计算之间的 grad norm 差异需要完成根因定位并形成数值验收结论。
  3. MTP 层当前选择跳过细粒度重计算,后续是否单独支持需要性能数据和生命周期验证。
  4. 是否将 CSA/MHC checkpoint 生命周期抽象成更通用的注册表或策略对象,以降低模型与 FB overlap 的耦合,需要后续评估。
  5. 是否增加运行时统计与 profiler 标记,用于输出每层释放量和重算耗时,需要社区确定默认行为。
  6. 非 Ascend、非融合 MHC 或其他 DeepSeek/CSA 变体的支持范围尚未定义。

附录

A. 主要代码位置

MindSpeed

MindSpeed-LLM

B. 术语表

术语 说明
CSA Compressed Sparse Attention,压缩稀疏注意力
MHC 模型中的 Hyper-Connection/混合连接模块,本文特指其 pre/post 路径
FB overlap MoE forward-backward overlap,前向、反向和通信交错调度
MTP Multi-Token Prediction
PP/VPP Pipeline Parallel / Virtual Pipeline Parallel
BDA Bias-Dropout-Add
CheckpointWithoutOutput 保留计算上下文但允许释放前向输出 storage 的细粒度 checkpoint
swap entry 一次 Layer 输入换出操作的独立跟踪对象,记录 D2H/H2D 状态及 CPU/NPU Tensor
安全挂钩点 已完成最后一次前向读取、且反向可通过该 Tensor hook 触发重算的图节点

C. 文档状态与后续补充

  • 参数说明已写入 docs/zh/pytorch/features/mcore/deepseek4_fine_grained_recompute.md,并加入特性列表与文档目录。
  • 性能验收完成后补充目标配置、峰值显存和吞吐实测结果。
  • grad norm 差异完成定位后补充数值一致性结论和推荐容差。
  • 若未来支持 MTP 或其他模型,更新适用范围和测试矩阵。

欢迎加入社区,感谢您对社区的贡献 🎉!

likedislike
wuweiqiang24wuweiqiang24成员
8月13日 添加了label:rfc
wuweiqiang24wuweiqiang24成员
8月13日 修改了issue 的描述
wuweiqiang24wuweiqiang24成员
8月13日 修改了issue 的描述
wuweiqiang24wuweiqiang24成员
8月14日 修改了issue 的描述
wuweiqiang24wuweiqiang24成员
8月14日 关联了pull request:feat: add mhc recompute
wuweiqiang24wuweiqiang24成员
3 天前 issue状态由 TODO 改变为 DONE
wuweiqiang24wuweiqiang24成员
3 天前 关闭了 issue
ascend-robotascend-robot成员
3 天前 添加了label:resolved