已合并
feat: add scheduled torch activation recomputation #1102
feat: add scheduled torch activation recomputation #1102
已合并
DavidFFFan创建于 8月1日
DavidFFFan
DavidFFFan成员
8月1日

What type of PR is this?

/kind feature


What does this PR do / why do we need it:

Torch 后端原先直接使用框架原生 non-reentrant checkpoint,无法获取 checkpoint frame,
因此不支持主动重计算以及跨 dx/dw GraphTask 复用重计算结果。

本 PR 在 HyperParallel Torch eager 后端补充 non-reentrant checkpoint 基础原语:

  • 提供 checkpoint handle 收集和 backward 前主动重计算能力;
  • 使用稳定 session id 缓存重计算结果,支持 dx/dw 独立调用时避免重复重计算;
  • 支持 retain_on_unpack 生命周期控制和幂等 session 清理;
  • 将 session activation 预绑定到 checkpoint frame,解决 NPU autograd worker 不继承主线程
    ContextVar 时无法识别 session 的问题;
  • saved-tensor unpack 热路径只读取 frame 字段,不执行 ContextVar 查询或全局加锁;
  • 支持 per-call early_stop、RNG/device/autocast 恢复、确定性检查和用户 context_fn;
  • eager 模式使用 HyperParallel 实现,compile 模式继续回退框架原生 checkpoint;
  • 普通 GraphTask nested checkpoint 保持可用,scheduled nested checkpoint 当前明确报错。

本次只交付底层主动重计算和 session 原语,不包含 PP stage 或调度器适配。
详细设计和后续实施边界见 RFC issue #310。


Which issue(s) this PR fixes:

Fixes #310


Test Plan and Test result:What scenarios were tested, and what were the verification results(Function, performance, reliability, etc.):

  • Core 与 Torch activation-checkpoint CPU UT:93 passed;
  • 单卡 Torch NPU 轻量 ST:5 passed;
  • CPU UT 覆盖多 checkpoint frame 共享 session、重复 iteration/session 隔离、
    early-stop、partial backward、异常清理和恢复;
  • NPU ST 覆盖真实 autograd worker 下的 dx/dw 分离与预重计算复用;
  • NPU ST 覆盖 dropout RNG mask、prefire 前后随机状态保持以及 dx/dw 数值一致;
  • NPU ST 覆盖在 autocast 外执行 prefire 时恢复 NPU bfloat16 autocast 配置;
  • py_compile、git diff --check 和 lizard 检查通过。

Self-checklist:(请自检,在[ ]内打上x,我们将检视你的完成情况,否则会导致pr无法合入)

likedislike
Pull Request已成功合入, 合并人@MindSpore-Bot
(感谢 DavidFFFan 的贡献)
DavidFFFanDavidFFFan成员
8月1日 创建了 pull request,commit 78f79e77
atomgit-bot
atomgit-bot
8月1日 评论:

变更摘要

本 PR 为 Hyper 框架新增"调度式激活重计算"能力:在 Torch 后端实现了 Hyper 自有的非可重入(non-reentrant)checkpoint 实现,支持 early_stop(在产生完 backward 所需张量后提前停止重计算)以及基于会话(session)的提前重计算调度——将一次重计算的结果按会话缓存,供 dx/dw 分离 backward 复用;同时将调度相关的 recompute_handle、recompute_session_ctx、clear_recompute_session 等接口暴露到平台层,并在 MindSpore 平台侧同步补充 session_id 非空校验。

主要改动

  • 新增 early_stop 参数:在 hyper_parallel/core/activation_checkpoint/activation_checkpoint.py 的 checkpoint() 中新增 early_stop(默认 True),进行 bool 类型校验并透传给平台实现;CheckpointWrapper 也将该配置透传至核心 checkpoint。
  • 新增 Torch 后端自研 checkpoint 实现:新增 hyper_parallel/platform/torch/activation_checkpoint/checkpoint.py,实现了 eager 非可重入 checkpoint(含 RNG 状态保持、autocast 上下文、determinism_check 元数据校验、context_fn 前向/重计算双上下文),编译场景回退到 PyTorch 原生 API,并新增 early_stop 关键字显式控制提前停止。
  • 新增会话调度式重计算接口:通过 recompute_handle_collector_ctx、recompute_handle、recompute_session_ctx、clear_recompute_session 及 _CheckpointFrame 实现按会话键缓存重计算结果,支持跨线程、跨 dx/dw 多次 backward 复用同一次预触发的重计算,并校验重计算张量数量与元数据一致性。
  • 平台层暴露新接口:hyper_parallel/platform/torch/platform.py 的 TorchPlatform.checkpoint 改为返回自研实现,并新增 recompute_handle_collector_ctx、recompute_handle、recompute_session_ctx、clear_recompute_session 静态方法;torch/activation_checkpoint/__init__.py 同步导出 CheckpointError、checkpoint、clear_recompute_session、recompute_handle、recompute_handle_collector_ctx、recompute_session_ctx。
  • 会话参数校验强化:hyper_parallel/platform/mindspore/platform.py 的 recompute_session_ctx 新增 session_id 非空校验,hyper_parallel/platform/platform.py 基类文档将 session_id 明确为必填且不可为 None,并规定上下文管理器 yield 该会话 id;Torch 实现亦对 session_id、retain_on_unpack、handle 等参数进行合法性校验。
  • 新增单元测试:新增 tests/ut/platform/torch/activation_checkpoint/test_checkpoint.py(覆盖梯度一致性、early_stop、会话共享重计算、跨线程复用、元数据不匹配报错等),并在现有 checkpoint 相关测试中补充 early_stop 透传与类型校验用例。
likedislike
不准确?
atomgit-bot
atomgit-bot
8月1日 评论:

代码审查

审查结论

已完成对全部 13 个变更文件的逐文件审查。

各文件审查结果

文件 结论
docs/api/api_reference.md 无问题:新增 early_stop/group_swap 参数文档与 Torch 2.6/2.7/2.9 行为说明,无危险操作或安全误导
docs/guide/activation_checkpoint.md 无问题:early_stop 用法示例与实际实现一致,无不安全命令
docs/guide/torch_non_reentrant_checkpoint_design.md 无问题:纯设计文档,无危险指令;但其 §15 关于 MindSpore 接收 early_stop 的断言与实现验证缺口相关(见 F1)
docs/index.md 无问题:仅新增设计文档链接
hyper_parallel/core/activation_checkpoint/activation_checkpoint.py F1 (P2):early_stop 无条件透传给 plat.checkpoint,MindSpore 后端 ms.recompute 契约未验证、无 MindSpore 侧改动/测试
hyper_parallel/platform/mindspore/platform.py 无问题:recompute_session_ctx 增加 None 校验,调用方均传非 None session id,安全
hyper_parallel/platform/platform.py 无问题:仅文档字符串更新
hyper_parallel/platform/torch/activation_checkpoint/init.py 无问题:导出新增符号,模块存在
hyper_parallel/platform/torch/activation_checkpoint/checkpoint.py F2 (P3):_StopRecomputationError 继承 Exception,可被用户 except Exception 吞掉破坏 early-stop。内核其余逻辑(generator 流程、holder/session 键控、RNG/autocast 恢复、嵌套 checkpoint 语义)经逐行核对为 torch 原生 non-reentrant 实现的忠实移植,未发现其他确定性缺陷
hyper_parallel/platform/torch/platform.py 无问题:checkpoint 属性及四个 session/collector 代理方法实现正确,核心仅由 core checkpoint 调用
tests/ut/core/activation_checkpoint/test_activation_checkpoint.py 无问题:新增用例断言与 core 行为一致
tests/ut/platform/torch/activation_checkpoint/test_checkpoint.py F3 (P3):缺失嵌套 checkpoint+stable session 与设备 RNG/autocast 高风险的测试覆盖
tests/ut/platform/torch/activation_checkpoint/test_checkpoint_wrapper.py 无问题:wrapper 透传 early_stop 的断言正确

统计

  • P0:0
  • P1:0
  • P2:1(MindSpore 后端 early_stop 跨平台契约风险)
  • P3:2(内部控制流异常可被吞掉;缺失高风险的嵌套/session 与设备 RNG/autocast 测试)

总体风险判断

本 PR 的核心(Torch eager non-reentrant checkpoint 内核)是一份结构严谨、与 PyTorch 原生实现高度对齐的移植,CPU 路径的数值、early-stop、session 预重算均有较完整单测支撑,未发现会导致梯度错误的确定性缺陷。主要风险集中在跨平台契约:公共 checkpoint() 签名变更(新增 early_stop 并在核心无条件透传)对 MindSpore 后端构成未经验证的破坏性变更,且该签名变更本身是公开 API 的行为变更(原会被转发给被包装函数的 early_stop 关键字现在被 checkpoint 消费),需要发行说明同步。综合判断:变更可接受,但建议在合入前确认 MindSpore 契约并补齐对应测试。

类型 数量
🔴 阻塞 0
🟡 建议 1

💬 仅评论

likedislike
不准确?
MindSpore-BotMindSpore-Bot成员
8月1日 添加了label:mindspore-cla/yes
司小南(机器人)司小南(机器人)成员
8月1日 添加了label:pr-check-pass
此处折叠了121条消息 查看更多
司小南(机器人)司小南(机器人)成员
8月6日 删除了label:ci-pipeline-running
司小南(机器人)司小南(机器人)成员
8月6日 添加了label:ci-pipeline-passed
Yyangzhenzhang成员
8月7日 通过审查
liuchongming74liuchongming74成员
8月7日 通过审查
MindSpore-BotMindSpore-Bot成员
8月7日 合入了pull request,合并节点 SHA:f832673dde8c20bef8d9a977a570ab9649852951