已合并
feat: add scheduled torch activation recomputation #1102
DavidFFFan创建于 8月1日
feat: add scheduled torch activation recomputation #1102
已合并
Pull Request已成功合入, 合并人@MindSpore-Bot
(感谢 DavidFFFan 的贡献)8月1日 创建了 pull request,commit 78f79e77
atomgit-bot
8月1日 评论:
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透传与类型校验用例。


不准确?
atomgit-bot
8月1日 评论:
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 |
💬 仅评论


不准确?
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
8月7日 通过审查
8月7日 通过审查
8月7日 合入了pull request,合并节点 SHA:f832673dde8c20bef8d9a977a570ab9649852951
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 基础原语:
retain_on_unpack生命周期控制和幂等 session 清理;ContextVar时无法识别 session 的问题;ContextVar查询或全局加锁;early_stop、RNG/device/autocast 恢复、确定性检查和用户context_fn;本次只交付底层主动重计算和 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.):
early-stop、partial backward、异常清理和恢复;
py_compile、git diff --check和 lizard 检查通过。Self-checklist:(请自检,在[ ]内打上x,我们将检视你的完成情况,否则会导致pr无法合入)