mindspore-ai / mindspore-ai/hyper-parallel
[RFC] checkpoint exclusion 边界激活优化:优先消除 RECOMPUTE→SAVE 冗余保存
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
建议按以下优先级演进:
- 优先优化
RECOMPUTE -> SAVE:避免 exclusion 区域的 backward save 固定住本可由
replay 重建的上游大激活; - 保持
SAVE -> RECOMPUTE默认缓存全部输出,延续 HyperParallel 当前简单、可靠的语义; - 若真实模型证明多输出或条件分支收益显著,再把
SAVE -> RECOMPUTE按需输出缓存设计为
opt-in 高级策略,不建议默认打开。
使用方法与设计对比
HyperParallel 可以在模型构造阶段声明哪些模块整体重算、哪些模块从重算中排除,不需要修改原有
construct/forward;torch_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 输出,而是:
- 为 SAVE 输出按 storage 注册弱引用 persist thunk;
- forward 进入显式
recompute=Trueregion 时,按输入 storage 查找 SAVE 生产者; - 只有被识别消费者使用的输出才写入生产者 durable output slot;
- 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 中重新生成。
当前行为可能同时付出:
- replay RMSNorm 的计算成本;
- 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
- wrapper 入口快照直接 tensor 输入,只持有 storage 弱引用和纯元数据;
- inner pack 先支持识别 exact input;
- pack 返回 invocation-scoped handle,不保存真实输入;
- replay 到 wrapper 时填充 slot;
- unpack 从 slot 读取;
- 不能证明安全时自动回退,不新增用户配置。
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,保持构造阶段声明式配置。
验证建议
RMSNorm(RECOMPUTE) -> GEMM(SAVE):- RMSNorm forward/replay 各执行一次;
- GEMM body 只在 forward 执行;
- GEMM backward 正确取得 replay 重建输入;
- value、dx、dw 与当前实现一致。
RMSNorm(SAVE) -> GEMM(RECOMPUTE):- 保持全输出缓存;
- 普通未包装 GEMM 无需额外标注即可 replay。
- SAVE 内部激活仍正常保存,不误判为输入。
- exact input、view input、多输入、kwargs、tuple/list 输出。
- in-place/version 变化安全回退。
SAVE -> SAVE输入不做 rederive;连续 SAVE 链只缓存最后一个输出。save_output=False覆盖单输出、多输出和嵌套输出,正向/replay Tensor 叶子与边界数量一致。- 同一 wrapper 多次调用、多个 invocation、连续 step。
- PP micro-batch、预触发 recompute、dx/dw 分离和 early-stop。
- 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
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- 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