mindspore-ai / mindspore-ai/hyper-parallel

[RFC] checkpoint exclusion 边界激活优化:优先消除 RECOMPUTE→SAVE 冗余保存

Open
#208 2 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
53
Forks
63
Avg merge
23h 45m
Merged PRs (30d)
63

Description

实施状态(2026-08)

方案已在 PR #1135 中实现,目标分支为
r1.0.0

  • RECOMPUTE -> SAVE:通过 invocation-scoped 延迟输入 handle 和输出边界 trigger,使 SAVE
    backward 所需的上游输入由 replay 重建,不再跨 forward/backward 常驻;
  • SAVE -> SAVE:新增 save_output 显式选项。默认 True 保持兼容;连续 SAVE 链的中间区域可设为
    False,只保留最后一个 SAVE 输出;
  • SAVE -> RECOMPUTE:继续使用默认 save_output=True 缓存真实输出,保证普通裸算子可以直接消费;
  • 区域内部的 saved-tensor pack/unpack 语义保持不变;本次没有引入隐式消费者推断,也不要求修改
    construct/forward

结论与建议

checkpoint_exclude_wrapper 已经支持在外层 checkpoint_wrapper 中排除一个
Cell/callable:forward 保存区域 backward 所需张量并缓存整个区域输出,replay 到达 wrapper
时直接返回缓存输出,跳过内部计算。基础设计见 #286。

参考 meta-pytorch/remat 的区域级 selective
rematerialization 设计,混合 RECOMPUTE 与 SAVE 区域时还存在两类边界激活问题:

RECOMPUTE -> SAVE
SAVE      -> RECOMPUTE

建议按以下优先级演进:

  1. 优先优化 RECOMPUTE -> SAVE:避免 exclusion 区域的 backward save 固定住本可由
    replay 重建的上游大激活;
  2. 保持 SAVE -> RECOMPUTE 默认缓存全部输出,延续 HyperParallel 当前简单、可靠的语义;
  3. 若真实模型证明多输出或条件分支收益显著,再把 SAVE -> RECOMPUTE 按需输出缓存设计为
    opt-in 高级策略,不建议默认打开。

使用方法与设计对比

HyperParallel 可以在模型构造阶段声明哪些模块整体重算、哪些模块从重算中排除,不需要修改原有
construct/forwardtorch_remat 则需要在 forward 的具体调用点插入 region 标记,属于侵入式改造。

对比项 HyperParallel torch_remat
外层重算单元 构造阶段声明:model.block = checkpoint_wrapper(model.block) 调用点包装:out = remat.checkpoint(region_name="block")(block)(x)
内部 SAVE 区域 构造阶段声明:block.gemm = checkpoint_exclude_wrapper(block.gemm) 必须修改 forward:x = remat.region(self.gemm, "gemm", recompute=False)(x)
内部 RECOMPUTE 区域 外层 checkpoint 内默认重算,不需要逐区域标记 裸调用默认重算;若要建立可识别的消费者边界,还需写 remat.region(fn, name, recompute=True)
原 forward/construct 无需修改,wrapper 对模型成员做声明式替换 需要修改,每个 SAVE/显式 RECOMPUTE 调用都要改写
区域身份 wrapper 实例与 invocation 内调用顺序 每个 checkpoint 内要求唯一字符串名称
SAVE 输出 当前缓存 wrapper 的完整输出 默认惰性注册;消费者触发后才持久化
裸消费者 输出已经缓存,通常无需额外处理 可能需要手动 recompute_needs_tensor
RECOMPUTE -> SAVE 输入 当前 inner hook 保存真实 tensor 可返回无 tensor handle,由 replay 重建后再供 unpack
当前后端 MindSpore PyNative PyTorch eager

HyperParallel 的完整声明式用法如下,Block.construct 本身无需增加 checkpoint 逻辑:

block.gemm = checkpoint_exclude_wrapper(block.gemm)
block = checkpoint_wrapper(block)

对应的 torch_remat 写法必须进入 forward 修改调用:

def forward(self, x):
    x = remat.region(self.norm, "norm", recompute=True)(x)
    x = remat.region(self.gemm, "gemm", recompute=False)(x)
    return x


out = remat.checkpoint(region_name="block")(block)(x)

从接入成本看,HyperParallel 当前接口更适合已有模型批量配置重算策略;torch_remat 控制更细,但要求
模型作者理解 region 边界、唯一名称、placeholder 和裸消费者。

边界一:SAVE -> RECOMPUTE

原理与 RMSNorm -> GEMM 场景

例如 RMSNorm 被 exclusion wrapper 排除重算,而后续 GEMM 仍由外层 checkpoint replay:

block.norm = checkpoint_exclude_wrapper(block.norm)  # SAVE
block = checkpoint_wrapper(block)


def construct(self, x):
    normalized = self.norm(x)  # replay 时跳过 body
    return self.gemm(normalized)  # replay 时重新执行

replay 时 RMSNorm 不执行,但 GEMM 必须拿到真实 normalized。因此 SAVE 区域输出必须跨越
forward/backward 保存:

forward:
  RMSNorm -> cache.save(wrapper_id, normalized)

replay:
  RMSNorm wrapper -> cache.pop(wrapper_id)
  GEMM -> 使用真实 normalized 重算

这首先是混合策略的正确性机制。HyperParallel 当前无条件缓存 exclusion wrapper 的整个输出,
因此无论下游是普通算子、Cell 还是自定义 kernel,都可以直接使用,模型作者不需要额外标注。

torch_remat 的按需输出优化

torch_remat 不立即保存全部 SAVE 输出,而是:

  1. 为 SAVE 输出按 storage 注册弱引用 persist thunk;
  2. forward 进入显式 recompute=True region 时,按输入 storage 查找 SAVE 生产者;
  3. 只有被识别消费者使用的输出才写入生产者 durable output slot;
  4. replay 跳过 SAVE body:有 slot 返回真实 tensor,无 slot 返回无数据 placeholder。

该方案在多输出、条件分支或 SAVE -> SAVE 中可以避免过度保存,但它无法自动观察裸算子。

需要手动指定保存的局限

外层 checkpoint 会让所有裸算子在 replay 中重新执行,但这不等于 remat 能观察裸算子的输入。例如:

def forward(self, x):
    normalized = remat.region(
        self.norm,
        "norm",
        recompute=False,
    )(x)
    return self.gemm(normalized)  # 普通裸调用

self.gemm(normalized) 确实会在外层 checkpoint replay 中再次运行,但因为它没有
remat.region 边界,remat 在 forward 时不知道 GEMM 将读取 SAVE 输出。replay 中 RMSNorm
返回 placeholder,GEMM 读取数据时才报错。

用户必须在第一个裸数据消费者之前手动指定保存:

remat.recompute_needs_tensor(normalized)
return self.gemm(normalized)

或者继续侵入 forward,把消费者也包装起来:

return remat.region(
    self.gemm,
    "gemm",
    recompute=True,
)(normalized)

该限制在以下常见代码中容易出现:

  • residual add、激活函数、contiguous() 等裸算子;
  • 未包装的 Module/Cell 调用;
  • fused/custom kernel;
  • SAVE 输出先经过产生新 storage 的裸算子,再进入显式 RECOMPUTE region。

只做 view/shape 等元数据操作时 placeholder 可以继续传播,但在第一个真正读取数据的裸算子前仍要
手动调用。漏标通常到 backward replay 才暴露,增加了接入和排错成本。

这也是不建议 HyperParallel 默认引入按需输出缓存的重要原因:HyperParallel 的完全声明式优势会被
recompute_needs_tensor 一类 forward 标注削弱。

何时按需输出缓存才有明显收益
  • SAVE 区域有多个较大的输出,只有部分进入 replay;
  • 条件分支只消费部分输出;
  • SAVE -> SAVE,中间输出不需要参与 replay;
  • 存在较大的 dead auxiliary output。

常见 RMSNorm、GEMM、AllToAll 通常只有一个主输出,而且下游必然消费;packed-token MoE 也常使用
单个 expert_out。这些场景按需保存的字节数接近缓存全部输出,但会增加 storage 索引、弱引用 thunk、
placeholder、手动标注和延迟错误。因此不引入运行时消费者推断,而采用显式、保守的 save_output

checkpoint_exclude_wrapper(module, save_output=True)   # 默认;输出可能被 RECOMPUTE/裸算子消费
checkpoint_exclude_wrapper(module, save_output=False)  # 仅用于输出整体直接传给下一个 SAVE 的中间区域

该选项只决定是否缓存区域输出,不改变区域内部 backward 数据的 pack。默认值保持原语义,模型作者只需在
可证明为连续 SAVE 的边界上 opt-in,不需要为每个消费者增加 region 或输出 mask。

边界补充:SAVE -> SAVE(显式融合,已实现)

目标场景与配置

连续 SAVE 区域都会在 replay 中跳过 body,因此中间区域的真实输出不需要跨 forward/backward 保存:

RECOMPUTE -> SAVE1(False) -> SAVE2(False) -> ... -> SAVEN(True)
block.save1 = checkpoint_exclude_wrapper(block.save1, save_output=False)
block.save2 = checkpoint_exclude_wrapper(block.save2, save_output=False)
block.save_last = checkpoint_exclude_wrapper(block.save_last)  # 默认 True

每个 SAVE 区域内部需要用于 backward 的 tensor 仍正常 pack;save_output=False 只消除区域输出缓存。
最后一个 SAVE 默认保存真实输出,以供 checkpoint 外部或后续 RECOMPUTE 区域使用。

replay placeholder 与边界数量

save_output=False 的 replay 不执行 SAVE body,仅返回零元素 placeholder。placeholder 的数据不会被后续
SAVE 使用,但其 Tensor 叶子数量必须与 forward 输出一致:每个正向 Tensor 输出都可能经过一个
_RecomputeBoundary,如果 replay 只返回单个 Tensor,多输出时会破坏 saved-tensor pack/unpack 顺序。

因此 forward 仅记录 output_tensor_count,不保存真实输出结构或 storage;replay 为每个 Tensor 叶子建立
对应的边界节点,底层复用同一个零元素 placeholder。这样保持激活序列一致,同时不重新引入输出显存占用。

只有 RECOMPUTE -> SAVE 的第一个 SAVE 需要输出边界来驱动上游 replay。由于 placeholder 本身不可微,
边界额外接收一个零元素、可微 trigger 以确保 MindSpore 建立自定义反向节点;其 backward 对 trigger 返回
None,不会产生或累计有效梯度。全局 trigger 只缓存一个零元素 Tensor,不持有业务 activation。

使用约束

save_output=False 只支持输出作为一个完整参数直接传给下一个 SAVE wrapper。两个 SAVE 之间不能有读取
数据的裸算子,也不能先拆包、索引或改变容器结构;存在此类消费时必须保留默认 save_output=True
该显式约束避免引入 output-use 自动推断,保持当前构造阶段声明式接口,不侵入模型源码。

边界二:RECOMPUTE -> SAVE(建议优先优化)

典型场景
RMSNorm (RECOMPUTE)
  -> normalized activation
  -> GEMM / fused expert / expensive communication region (SAVE)

GEMM 被 exclusion wrapper 排除重算,但 GEMM backward 计算 dW 通常需要保存输入
normalized activation。当前 inner saved-tensor hook 会保留该真实输入直到 backward。

问题是 normalized activation 来自 RECOMPUTE 区域,本来就会在 checkpoint replay 中重新生成。
当前行为可能同时付出:

  1. replay RMSNorm 的计算成本;
  2. normalized activation 跨 forward/backward 常驻的显存成本。

该激活一般为 [tokens, hidden],每层都会出现;MoE 中还可能是 dispatch 后送入 grouped GEMM 的
packed token activation,累计占用比多输出辅助张量更值得优先优化。

优化原理

在 exclusion wrapper 的 saved-tensor hook 中区分“区域输入”和“区域内部产生的激活”:

forward:
  1. wrapper 入口记录 tensor 输入:
     path + weak storage identity + dtype/shape/stride/offset + version
  2. inner pack 收到 backward saved tensor
  3. 若它等于某个可重算输入或其可重建 view:
       不保存真实 tensor
       返回 SavedInputRef(slot_name)
       记录 SavedInputRecipe(path, view metadata)
  4. 区域内部产生的 saved tensor 仍按现有方式保存

replay:
  5. 上游 RECOMPUTE 重新生成 wrapper 输入
  6. exclusion wrapper 跳过 body,但在返回缓存输出前:
       根据 recipe 从当前 args/kwargs 取出输入
       填充 invocation-scoped rederived slot

backward:
  7. inner unpack 根据 SavedInputRef.slot_name 返回 replay 重建的 tensor

handle/recipe 只保存路径和布局等元数据,不持有 tensor 或设备 storage。以下情况应保守回退到当前真实
保存路径:

  • 输入来自另一个 exclusion/SAVE 区域,replay 不会重建;
  • exclusion body 内发生 in-place mutation,version 已变化;
  • dtype reinterpret、不安全 view 或无法稳定重建的容器;
  • replay 路径、输入布局或 wrapper 调用顺序与 forward 不一致。

saved-input unpack 必须晚于 replay 到达对应 exclusion wrapper。需要验证普通 non-reentrant
checkpoint、PP 预触发和 dx/dw 分离下的执行顺序;若无法天然保证,应通过 checkpoint 输出边界 trigger
或 MindSpore recompute handle 先驱动 replay,再允许内部 backward unpack。

为什么优先级更高

该优化:

  • 命中常见的单输入 GEMM、Linear、expert GEMM,不依赖多输出或条件分支;
  • 直接消除 [tokens, hidden] 级别的跨 backward 常驻激活;
  • 对用户透明,不需要修改原 construct/forward;
  • 不改变 checkpoint_wrapper/checkpoint_exclude_wrapper 声明式用法;
  • 当前全量缓存 SAVE 输出仍可保证 SAVE -> RECOMPUTE 正确。

建议实施顺序

P0:透明实现 RECOMPUTE -> SAVE saved-input rederive
  1. wrapper 入口快照直接 tensor 输入,只持有 storage 弱引用和纯元数据;
  2. inner pack 先支持识别 exact input;
  3. pack 返回 invocation-scoped handle,不保存真实输入;
  4. replay 到 wrapper 时填充 slot;
  5. unpack 从 slot 读取;
  6. 不能证明安全时自动回退,不新增用户配置。
P1:补充 view、诊断和内存报告
  • 支持可验证的 input view 重建;
  • 检查 dtype、layout、version;
  • slot 缺失时错误包含 wrapper、输入路径和 phase;
  • resident SAVE activation 与 replay-rederived input 分开统计。
P2:以显式 save_output 优化 SAVE -> SAVE(已实现)
  • 默认 save_output=True,不改变 SAVE -> RECOMPUTE 和裸消费者语义;
  • 连续 SAVE 链的中间区域显式配置 False,最后一个区域保持 True
  • forward 记录 Tensor 输出叶子数量,replay 用零元素 placeholder 复现相同边界序列;
  • 不做隐式消费者识别,不引入 recompute_needs_tensor,保持构造阶段声明式配置。

验证建议

  1. RMSNorm(RECOMPUTE) -> GEMM(SAVE)
    • RMSNorm forward/replay 各执行一次;
    • GEMM body 只在 forward 执行;
    • GEMM backward 正确取得 replay 重建输入;
    • value、dx、dw 与当前实现一致。
  2. RMSNorm(SAVE) -> GEMM(RECOMPUTE)
    • 保持全输出缓存;
    • 普通未包装 GEMM 无需额外标注即可 replay。
  3. SAVE 内部激活仍正常保存,不误判为输入。
  4. exact input、view input、多输入、kwargs、tuple/list 输出。
  5. in-place/version 变化安全回退。
  6. SAVE -> SAVE 输入不做 rederive;连续 SAVE 链只缓存最后一个输出。
  7. save_output=False 覆盖单输出、多输出和嵌套输出,正向/replay Tensor 叶子与边界数量一致。
  8. 同一 wrapper 多次调用、多个 invocation、连续 step。
  9. PP micro-batch、预触发 recompute、dx/dw 分离和 early-stop。
  10. MoE 增加 dispatch(RECOMPUTE) -> grouped GEMM(SAVE) 端到端用例。

显存验证应确认 [tokens, hidden] 输入在 forward 后不再由 exclusion saved-tensor payload 持有,
并分别报告 resident bytes、replay 临时峰值、峰值 HBM、step time 和 Host 开销。

当前验证结果

  • MindSpore activation checkpoint UT:59 passed,1 skipped;Core activation checkpoint UT:23 passed;
  • PR changed 门禁:Core UT 23 passed、MindSpore UT 22 passed;
  • 单卡连续 SAVE 链 ST 通过;activation checkpoint 并行组 8 个子场景全部通过;
  • RMSNorm→MatMul 20 层用例覆盖功能、输出显存和 Host 性能路径;
  • MindSpore 2.10.0 长稳:全局缓存 trigger 运行 50 万步,Device allocated 始终 512 B、reserved
    始终 2 MiB,10 万至 50 万步 RSS 仅增加 76 KiB;
  • 每次新建 trigger 的对照运行 25 万步,Device 指标一致,10 万至 25 万步 RSS 仅增加 44 KiB;
  • 两条曲线均在预热后进入平台期,trigger 引用计数保持 3、.grad 始终为 None,未观察到
    lru_cache 引入的持续 Device 或 Host 内存泄漏。

验收标准

  • 默认 API、声明式用法和当前 SAVE -> RECOMPUTE 行为不变;
  • 可重算的 exclusion saved input 不跨 forward/backward 常驻;
  • 不可安全重建的输入保守回退,不牺牲精度;
  • 单卡、连续 step、PP dx/dw 和 MoE 场景无 cache/slot 串用;
  • benchmark 能量化减少的 resident bytes 和新增 replay 临时峰值;
  • 文档明确两类边界的语义、收益场景和限制。

schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 307
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/307

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Review checkpoint_wrapper and checkpoint_exclude_wrapper first, then inspect the activation checkpoint and Core activation checkpoint tests described in the issue. Compare the implementation in PR #1135 with the listed RECOMPUTE→SAVE, SAVE→RECOMPUTE, and SAVE→SAVE cases. Done means safe saved-input rederivation with conservative fallback, unchanged default behavior, passing tests, and measured memory results for the listed scenarios.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
backend, machine-learning, performance
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
18/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.