flowchart TD
A["Checkpoint forward"] --> B["Collector records frame handle"]
B --> C["Prefire with stable session"]
C --> D["Replay and cache saved tensors on frame"]
D --> E["dx GraphTask: retain_on_unpack=True"]
E --> F["dw GraphTask: retain_on_unpack=False"]
F --> G["clear_recompute_session"]
D -. "no session" .-> H["GraphTask id fallback"]
6.2 架构参考
flowchart LR
subgraph Core["Core API"]
Checkpoint["checkpoint / checkpoint_wrapper"]
State["Recompute invocation context"]
end
subgraph TorchBackend["Torch Backend"]
Frame["CheckpointFrame"]
Hooks["saved_tensors_hooks"]
Session["Session control plane"]
Cache["Per-session recomputed tensors"]
end
subgraph Consumers["Consumers"]
Prefire["Prefire scheduler"]
DX["dx GraphTask"]
DW["dw GraphTask"]
end
Checkpoint --> Frame
State --> Hooks
Frame --> Hooks
Session --> Frame
Frame --> Cache
Prefire --> Session
DX --> Cache
DW --> Cache
6.3 时序参考
sequenceDiagram
participant U as Upper Scheduler
participant C as Checkpoint Frame
participant W as NPU Autograd Worker
participant D as dx/dw Consumer
U->>C: collect handle during forward
U->>C: recompute_handle(handle, session)
C->>C: bind SessionActivation to frame
C->>C: replay and cache saved tensors
D->>W: launch dx GraphTask
W->>C: unpack reads frame.active_session
C-->>W: retained recomputed tensor
D->>W: launch dw GraphTask
W->>C: unpack reads same session
C-->>W: final recomputed tensor
U->>C: clear_recompute_session(session)
背景
MindSpore non-reentrant checkpoint 已具备重计算 handle 收集、backward 前主动重计算、稳定 session 缓存以及跨 dx/dw 复用重计算结果的能力。Torch 原生
use_reentrant=Falsecheckpoint 只在 saved tensor 首次 unpack 时按 GraphTask 懒触发重计算,不公开 checkpoint frame,也不支持由上层调度器主动触发和跨 GraphTask 复用。在 dx/dw 分离场景中,输入梯度和权重梯度由不同 autograd GraphTask 计算。如果直接使用 Torch 原生实现,两次 backward 会分别执行 checkpoint replay,增加计算开销。NPU autograd 还可能在独立 worker 线程执行 saved-tensor unpack,主线程设置的
ContextVar不会自动传播,不能单纯依赖动态上下文识别 session。本 RFC 用于说明 Torch 后端补齐上述底层原语的设计。当前只交付 checkpoint 内核与 session 生命周期能力,PP stage 和调度器接入后续单独设计。
1. 基本信息
core/activation_checkpoint、platform/torch/activation_checkpoint2. 背景
use_reentrant=False语义上扩展调度基础原语,不修改或 patch 框架源码3. 目标和非目标
3.1 目标
retain_on_unpack和幂等清理原语。ContextVar时仍能识别 frame 对应的 session。ContextVar查询。early_stop、RNG、device/autocast 恢复、确定性检查、kwargs 和context_fn组合。3.2 非目标
checkpoint_sequential。torch.compile内部 checkpoint,compile 模式回退 Torch 原生实现。checkpoint_exclude_wrapper作为独立能力后续适配;当前 invocation state/session 方案与其兼容。4. 相关实现参考
use_reentrant=False5. 对外接口
5.1 接口定义
with platform.recompute_handle_collector_ctx() as handles: output = checkpoint(function, *args, early_stop=True, **kwargs) session_id = ("micro_batch", micro_batch_id) try: with platform.recompute_session_ctx(session_id, retain_on_unpack=True): for handle in handles: platform.recompute_handle(handle, session_id) with platform.recompute_session_ctx(session_id, retain_on_unpack=True): dx = compute_input_grad(output) with platform.recompute_session_ctx(session_id, retain_on_unpack=False): dw = compute_weight_grad(output) finally: platform.clear_recompute_session(session_id)early_stopboolTrueTrue/FalseValueErrorsession_idNone且可哈希ValueErrorretain_on_unpackboolFalseTrue/FalseValueErrorhandleValueError5.2 使用示例
上层使用顺序固定为:收集 handle → prefire → dx 保留 → dw 最终消费 → finally 清理。即使最终消费者已完成,也必须执行
clear_recompute_session(),释放 partial backward 或未消费 holder 关联的数据。5.3 接口说明
接口与 MindSpore 已有平台抽象保持一致。
session_id必须由调用方显式提供,避免匿名 id 无法传递到后续 dx/dw。handle 对用户保持不透明,上层不依赖 Torch checkpoint 私有 frame 类型。6. 方案设计
6.1 总体流程
flowchart TD A["Checkpoint forward"] --> B["Collector records frame handle"] B --> C["Prefire with stable session"] C --> D["Replay and cache saved tensors on frame"] D --> E["dx GraphTask: retain_on_unpack=True"] E --> F["dw GraphTask: retain_on_unpack=False"] F --> G["clear_recompute_session"] D -. "no session" .-> H["GraphTask id fallback"]6.2 架构参考
flowchart LR subgraph Core["Core API"] Checkpoint["checkpoint / checkpoint_wrapper"] State["Recompute invocation context"] end subgraph TorchBackend["Torch Backend"] Frame["CheckpointFrame"] Hooks["saved_tensors_hooks"] Session["Session control plane"] Cache["Per-session recomputed tensors"] end subgraph Consumers["Consumers"] Prefire["Prefire scheduler"] DX["dx GraphTask"] DW["dw GraphTask"] end Checkpoint --> Frame State --> Hooks Frame --> Hooks Session --> Frame Frame --> Cache Prefire --> Session DX --> Cache DW --> Cache6.3 时序参考
sequenceDiagram participant U as Upper Scheduler participant C as Checkpoint Frame participant W as NPU Autograd Worker participant D as dx/dw Consumer U->>C: collect handle during forward U->>C: recompute_handle(handle, session) C->>C: bind SessionActivation to frame C->>C: replay and cache saved tensors D->>W: launch dx GraphTask W->>C: unpack reads frame.active_session C-->>W: retained recomputed tensor D->>W: launch dw GraphTask W->>C: unpack reads same session C-->>W: final recomputed tensor U->>C: clear_recompute_session(session)6.4 关键逻辑
_CheckpointFrame管理 forward holder、per-session replay tensor、metadata 和active_session。RLock保护。frame.active_session;没有 activation 时使用当前 GraphTask id。ContextVar,用于生命周期校验和 scheduled nested 检测,不依赖线程传播。retain_on_unpack=True允许多个消费者复用;最终消费者设为False,finally 继续执行幂等 clear。early_stop=True在所有 forward holder 对应 tensor 已产生后抛内部控制流异常结束 replay。6.5 代码改动点
core/activation_checkpointearly_stop和 recompute context 组合platform/platform.pyplatform/torch/activation_checkpointplatform/torch/platform.pyplatform/mindspore/platform.pytests6.6 方案取舍
ContextVar该方案的主要代价是 HyperParallel 需要维护一份 Torch eager non-reentrant checkpoint 内核,并在 Torch 升级时持续对照原生实现。
7. 组件依赖
8. 约束与兼容性
set_checkpoint_early_stop,使用 Hyper per-callearly_stop;compile 回退原生约束:调用方必须遵循“prefire 完成 → dx/dw 有序消费 → finally clear”,不得在同一 frame 上并发激活不同 session。
9. 验证设计
9.1 测试范围
9.2 用例分层
9.3 交互验证
9.4 性能 / 显存验证
本期以功能和调用次数验收,不新增性能门禁;stage 接入后补充端到端 step time 和峰值显存数据。
10. 实现计划
checkpoint_exclude_wrapper适配