已合并
feat: add scheduled torch activation recomputation #1102
feat: add scheduled torch activation recomputation #1102
已合并
DavidFFFan创建于 8月1日
共 13 个文件变更+1704-11
@@ -845,6 +845,8 @@ checkpoint(
845 swap_inputs: bool = False,845 swap_inputs: bool = False,
846 policy_fn: Optional[Callable] = None,846 policy_fn: Optional[Callable] = None,
847 context_fn: Optional[Callable[[], Tuple[object, object]]] = None,847 context_fn: Optional[Callable[[], Tuple[object, object]]] = None,
848+ group_swap: bool = False,
849+ early_stop: bool = True,
848 **kwargs,850 **kwargs,
849)851)
850```852```
@@ -857,6 +859,12 @@ checkpoint(
857| `swap_inputs` | `bool` | `False` | 是否将 checkpoint 保存的输入 offload 到 CPU |859| `swap_inputs` | `bool` | `False` | 是否将 checkpoint 保存的输入 offload 到 CPU |
858| `policy_fn` | `callable` | `None` | SAC 逐算子策略:`(ctx, op, *args, **kwargs) -> CheckpointPolicy` |860| `policy_fn` | `callable` | `None` | SAC 逐算子策略:`(ctx, op, *args, **kwargs) -> CheckpointPolicy` |
859| `context_fn` | `callable` | `None` | 返回 `(forward_ctx, recompute_ctx)` 的无参工厂;可与 `policy_fn` 组合 |861| `context_fn` | `callable` | `None` | 返回 `(forward_ctx, recompute_ctx)` 的无参工厂;可与 `policy_fn` 组合 |
862+| `group_swap` | `bool` | `False` | 是否对 `MUST_SWAP` tensor 启用分组 copy 融合 |
863+| `early_stop` | `bool` | `True` | 是否在产生全部 backward 所需 tensor 后提前停止重计算 |
864+ 
865+Torch 2.6、2.7、2.9 eager 模式统一使用 HyperParallel non-reentrant 实现。`early_stop` 只接受单次调用
866+关键字或 `checkpoint_kwargs` 配置,不继承外层 `torch.utils.checkpoint.set_checkpoint_early_stop()`。
867+Torch backend 还会消费 `preserve_rng_state`、`determinism_check` 和 `debug` 等 checkpoint 保留关键字。
860 868 
861---869---
862 870 
@@ -890,7 +898,7 @@ swap(
890checkpoint_wrapper(module, **checkpoint_kwargs) -> CheckpointWrapper898checkpoint_wrapper(module, **checkpoint_kwargs) -> CheckpointWrapper
891```899```
892 900 
893-`checkpoint_kwargs` 与 `checkpoint` 一致,例如 `policy_fn`、`swap_inputs`。901+`checkpoint_kwargs` 与 `checkpoint` 一致,例如 `policy_fn`、`swap_inputs`、`early_stop`。
894 902 
895---903---
896 904 
@@ -53,8 +53,24 @@ def context_fn():
53 return forward_context(), recompute_context()53 return forward_context(), recompute_context()
54 54 
55output = checkpoint(model.layer, x, context_fn=context_fn)55output = checkpoint(model.layer, x, context_fn=context_fn)
56+ 
57+# 控制是否在产生全部 backward 所需 tensor 后提前停止重计算
58+output = checkpoint(model.layer, x, early_stop=True) # 默认值,减少无效尾部计算
59+output = checkpoint(model.layer, x, early_stop=False) # 完整执行整个重计算区域
56```60```
57 61 
62+Torch eager 模式在 2.6、2.7 和 2.9 上统一使用 HyperParallel 的 non-reentrant 实现,并固定
63+`use_reentrant=False`。`early_stop` 只通过 Hyper 的单次调用关键字或配置字典控制:
64+ 
65+```python
66+checkpoint_kwargs = {"early_stop": False, "preserve_rng_state": True}
67+output = checkpoint(model.layer, x, **checkpoint_kwargs)
68+```
69+ 
70+外层的 `torch.utils.checkpoint.set_checkpoint_early_stop()` 不会覆盖 Hyper eager checkpoint 的参数。
71+`checkpoint_wrapper(model.layer, early_stop=False)` 使用相同规则。图编译状态暂时回退到 Torch 原生
72+non-reentrant checkpoint;完整的 compile/context_fn 组合适配不在当前版本范围内。
73+ 
58### 2. 函数式 swap74### 2. 函数式 swap
59 75 
60```python76```python
@@ -108,6 +108,7 @@ def checkpoint(
108 policy_fn: Optional[Callable] = None,108 policy_fn: Optional[Callable] = None,
109 context_fn: Optional[Callable[[], Tuple[object, object]]] = None,109 context_fn: Optional[Callable[[], Tuple[object, object]]] = None,
110 group_swap: bool = False,110 group_swap: bool = False,
111+ early_stop: bool = True,
111 **kwargs,112 **kwargs,
112):113):
113 """114 """
@@ -129,11 +130,17 @@ def checkpoint(
129 order and exit in reverse.130 order and exit in reverse.
130 group_swap (bool, optional): Whether MUST_SWAP tensors participate in group copy fusion.131 group_swap (bool, optional): Whether MUST_SWAP tensors participate in group copy fusion.
131 Only effective when ``policy_fn`` is provided. Default: ``False``.132 Only effective when ``policy_fn`` is provided. Default: ``False``.
133+ early_stop (bool, optional): Whether recomputation stops after all tensors needed by
134+ backward have been produced. This per-call keyword is the only supported way to
135+ configure early stop. Default: ``True``.
132 **kwargs: Additional keyword arguments to pass to the function.136 **kwargs: Additional keyword arguments to pass to the function.
133 137 
134 Returns:138 Returns:
135 The result of applying the function with checkpointing.139 The result of applying the function with checkpointing.
136 """140 """
141+ if not isinstance(early_stop, bool):
142+ raise ValueError(f"early_stop must be bool, but got {type(early_stop).__name__}.")
143+ 
137 factories: list = [create_recompute_contexts]144 factories: list = [create_recompute_contexts]
138 if policy_fn is not None:145 if policy_fn is not None:
139 factories.append(partial(plat.create_selective_checkpoint_contexts, policy_fn, group_swap=group_swap))146 factories.append(partial(plat.create_selective_checkpoint_contexts, policy_fn, group_swap=group_swap))
@@ -148,7 +155,12 @@ def checkpoint(
148 context = partial(plat.async_save_on_cpu, group_swap=group_swap) if swap_inputs else contextlib.nullcontext155 context = partial(plat.async_save_on_cpu, group_swap=group_swap) if swap_inputs else contextlib.nullcontext
149 with context():156 with context():
150 return plat.checkpoint(157 return plat.checkpoint(
151- function, *args, context_fn=composed_context_fn, use_reentrant=False, **kwargs158+ function,
159+ *args,
160+ context_fn=composed_context_fn,
161+ use_reentrant=False,
162+ early_stop=early_stop,
163+ **kwargs,
152 )164 )
153 165 
154 166 
@@ -440,7 +440,7 @@ class _MSAsyncA2ALazyBwd(_Function):
440 return AsyncCollectiveTensor(actual_output, work)440 return AsyncCollectiveTensor(actual_output, work)
441 441 
442 @staticmethod442 @staticmethod
443- def backward(ctx, grad_output):443+ def backward(ctx, grad_output): # pylint: disable=arguments-differ
444 """Symmetric reverse a2a; returns :class:`AsyncCollectiveTensor`."""444 """Symmetric reverse a2a; returns :class:`AsyncCollectiveTensor`."""
445 # If grad_output is still lazy, force unwrap before issuing the445 # If grad_output is still lazy, force unwrap before issuing the
446 # reverse a2a (which is itself a "real" op on the data).446 # reverse a2a (which is itself a "real" op on the data).
@@ -612,7 +612,7 @@ class _MSSyncHookFunction(_Function):
612 return _MSSyncHookFunction._passthrough(x)612 return _MSSyncHookFunction._passthrough(x)
613 613 
614 @staticmethod614 @staticmethod
615- def backward(ctx, grad_output):615+ def backward(ctx, grad_output): # pylint: disable=arguments-differ
616 """Mirror of :meth:`forward` using ``_BWD_ROLES``."""616 """Mirror of :meth:`forward` using ``_BWD_ROLES``."""
617 hook_name = ctx.hook_name617 hook_name = ctx.hook_name
618 coordinator = ctx.coordinator618 coordinator = ctx.coordinator
@@ -657,7 +657,7 @@ class _MSAsyncA2AFunction(_Function):
657 return _a2a_reconstruct_ms(out_perm, concat_dim)657 return _a2a_reconstruct_ms(out_perm, concat_dim)
658 658 
659 @staticmethod659 @staticmethod
660- def backward(ctx, grad_output):660+ def backward(ctx, grad_output): # pylint: disable=arguments-differ
661 """Launch async head->seq A2A for backward overlap, or return zero grad."""661 """Launch async head->seq A2A for backward overlap, or return zero grad."""
662 if ctx.handle_box is not None:662 if ctx.handle_box is not None:
663 g = grad_output.contiguous()663 g = grad_output.contiguous()
@@ -695,7 +695,7 @@ class _MSAsyncAllGatherFunction(_Function):
695 return _move_dim_from_front(out_perm, gather_dim)695 return _move_dim_from_front(out_perm, gather_dim)
696 696 
697 @staticmethod697 @staticmethod
698- def backward(ctx, grad_output):698+ def backward(ctx, grad_output): # pylint: disable=arguments-differ
699 """Launch reverse reduce-scatter for the all-gather."""699 """Launch reverse reduce-scatter for the all-gather."""
700 grad_perm = _move_dim_to_front(grad_output.contiguous(), ctx.gather_dim)700 grad_perm = _move_dim_to_front(grad_output.contiguous(), ctx.gather_dim)
701 output_shape = list(grad_perm.shape)701 output_shape = list(grad_perm.shape)
@@ -1862,6 +1862,8 @@ class MindSporePlatform(Platform):
1862 1862 
1863 @staticmethod1863 @staticmethod
1864 def recompute_session_ctx(session_id, retain_on_unpack=False):1864 def recompute_session_ctx(session_id, retain_on_unpack=False):
1865+ if session_id is None:
1866+ raise ValueError("session_id must not be None.")
1865 # pylint: disable=C04151867 # pylint: disable=C0415
1866 from mindspore.common.recompute import _recompute_session_ctx1868 from mindspore.common.recompute import _recompute_session_ctx
1867 return _recompute_session_ctx(session_id=session_id, retain_on_unpack=retain_on_unpack)1869 return _recompute_session_ctx(session_id=session_id, retain_on_unpack=retain_on_unpack)
@@ -13,6 +13,9 @@
13# limitations under the License.13# limitations under the License.
14# ============================================================================14# ============================================================================
15"""framework platform api"""15"""framework platform api"""
16+# Backend platform modules intentionally import this abstraction to register
17+# their implementations; the resulting import cycle is architectural.
18+# pylint: disable=cyclic-import
16import os19import os
17from datetime import timedelta20from datetime import timedelta
18from enum import auto, Enum21from enum import auto, Enum
@@ -1548,15 +1551,19 @@ class Platform:
1548 """Context manager binding recompute unpack to a caller-provided session.1551 """Context manager binding recompute unpack to a caller-provided session.
1549 1552 
1550 Args:1553 Args:
1551- session_id: Stable session key. Recompute caches are keyed by this1554+ session_id: Required stable session key. Recompute caches are keyed
1552- instead of the transient autodiff engine id, so a re-run fired1555+ by this instead of the transient autodiff engine id, so a re-run
1553- under one engine can be reused by another.1556+ fired under one engine can be reused by another. Must not be
1557+ ``None``.
1554 retain_on_unpack (bool): When ``True``, unpack returns recomputed1558 retain_on_unpack (bool): When ``True``, unpack returns recomputed
1555 tensors without popping them, so a later backward can consume1559 tensors without popping them, so a later backward can consume
1556 them. Default: ``False``.1560 them. Default: ``False``.
1557 1561 
1558 Returns:1562 Returns:
1559 A context manager activating the session for its scope.1563 A context manager activating the session for its scope.
1564+ 
1565+ Yields:
1566+ The supplied session id.
1560 """1567 """
1561 raise NotImplementedError("Platform subclasses must implement recompute_session_ctx")1568 raise NotImplementedError("Platform subclasses must implement recompute_session_ctx")
1562 1569 
@@ -15,10 +15,24 @@
15"""Activation checkpointing related interfaces"""15"""Activation checkpointing related interfaces"""
16from .checkpoint_wrapper import CheckpointWrapper, ckpt_wrapper16from .checkpoint_wrapper import CheckpointWrapper, ckpt_wrapper
17from .activation_swap import swap_wrapper, swap_tensor_wrapper17from .activation_swap import swap_wrapper, swap_tensor_wrapper
18+from .checkpoint import (
19+ CheckpointError,
20+ checkpoint,
21+ clear_recompute_session,
22+ recompute_handle,
23+ recompute_handle_collector_ctx,
24+ recompute_session_ctx,
25+)
18 26 
19__all__ = [27__all__ = [
20 "CheckpointWrapper",28 "CheckpointWrapper",
21 "ckpt_wrapper",29 "ckpt_wrapper",
22 "swap_wrapper",30 "swap_wrapper",
23 "swap_tensor_wrapper",31 "swap_tensor_wrapper",
32+ "CheckpointError",
33+ "checkpoint",
34+ "clear_recompute_session",
35+ "recompute_handle",
36+ "recompute_handle_collector_ctx",
37+ "recompute_session_ctx",
24]38]
@@ -0,0 +1,661 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""Eager non-reentrant activation checkpointing for the Torch backend.
16+ 
17+The saved-tensor hook algorithm is adapted from ``torch.utils.checkpoint`` in
18+PyTorch release/2.9. Hyper owns the scheduling and session extensions here so
19+the implementation can run consistently on PyTorch 2.6, 2.7, and 2.9 without
20+patching the installed framework.
21+"""
22+import contextlib
23+import contextvars
24+import threading
25+import uuid
26+import warnings
27+import weakref
28+from collections import defaultdict
29+from typing import Any, Callable, DefaultDict, Dict, Generator, Iterator, List, Optional, Tuple
30+ 
31+import torch
32+from torch.utils._pytree import tree_map
33+from torch.utils.checkpoint import DefaultDeviceType
34+from torch.utils.checkpoint import checkpoint as torch_checkpoint
35+from torch.utils.checkpoint import set_checkpoint_early_stop
36+ 
37+ 
38+_DEFAULT_DETERMINISM_MODE = "default"
39+_RECOMPUTE_COLLECTOR = contextvars.ContextVar("hyper_recompute_collector", default=None)
40+_RECOMPUTE_SESSION = contextvars.ContextVar("hyper_recompute_session", default=None)
41+_SESSION_FRAMES: DefaultDict[Any, weakref.WeakSet] = defaultdict(weakref.WeakSet)
42+_SESSION_FRAMES_LOCK = threading.RLock()
43+ 
44+ 
45+class CheckpointError(RuntimeError):
46+ """Raised when checkpoint forward and recomputation are inconsistent."""
47+ 
48+ 
49+class _Handle:
50+ """Identity key for one recomputed saved tensor."""
51+ 
52+ 
53+class _Holder:
54+ """Saved-tensor placeholder containing handles keyed by recompute session."""
55+ 
56+ def __init__(self) -> None:
57+ """Initialize an empty per-session handle mapping."""
58+ self.handles: Dict[Any, Optional[_Handle]] = {}
59+ 
60+ 
61+class _StopRecomputationError(Exception):
62+ """Internal control-flow exception used by early-stop recomputation."""
63+ 
64+ 
65+class _SessionActivation:
66+ """Control-plane state shared by checkpoint frames in one session scope."""
67+ 
68+ def __init__(self, session_id: Any, retain_on_unpack: bool) -> None:
69+ """Initialize one scoped session activation."""
70+ self.session_id = session_id
71+ self.retain_on_unpack = retain_on_unpack
72+ self.frames: weakref.WeakSet = weakref.WeakSet()
73+ 
74+ 
75+class _NoopSaveInputs(torch.autograd.Function):
76+ """Save checkpoint inputs without adding a meaningful forward operation."""
77+ 
78+ @staticmethod
79+ def forward(*args: Any) -> Any:
80+ """Return a dummy output whose grad function retains checkpoint inputs."""
81+ del args
82+ return torch.empty((0,))
83+ 
84+ @staticmethod
85+ def setup_context(ctx: Any, inputs: Tuple[Any, ...], output: Any) -> None:
86+ """Save tensor inputs while retaining non-tensor input structure."""
87+ del output
88+ tensor_pairs = [(index, value) for index, value in enumerate(inputs) if isinstance(value, torch.Tensor)]
89+ tensor_indices, tensors = zip(*tensor_pairs)
90+ index_to_saved_index = {input_index: saved_index for saved_index, input_index in enumerate(tensor_indices)}
91+ stored_args = [None if isinstance(value, torch.Tensor) else value for value in inputs]
92+ 
93+ def get_args(saved_tensors: Tuple[Any, ...]) -> List[Any]:
94+ """Reconstruct the original checkpoint arguments."""
95+ restored_args = [
96+ saved_tensors[index_to_saved_index[index]] if index in tensor_indices else value
97+ for index, value in enumerate(stored_args)
98+ ]
99+ return restored_args[1:]
100+ 
101+ ctx.get_args = get_args
102+ ctx.save_for_backward(*tensors)
103+ 
104+ @staticmethod
105+ def backward(ctx: Any, *grad_outputs: Any) -> None:
106+ """Reject direct backward through the internal input saver."""
107+ del ctx, grad_outputs
108+ raise CheckpointError("The internal checkpoint input saver must not be backwarded directly.")
109+ 
110+ 
111+class _CheckpointFrame:
112+ """State shared by one checkpoint forward and its recomputations."""
113+ 
114+ def __init__(self, recompute_fn: Callable, early_stop: bool, metadata_fn: Optional[Callable]) -> None:
115+ """Initialize frame state captured by the saved-tensor hooks."""
116+ self.recompute_fn = recompute_fn
117+ self.input_saver = None
118+ self.weak_holders: List[weakref.ReferenceType] = []
119+ self.recomputed: DefaultDict[Any, weakref.WeakKeyDictionary] = defaultdict(weakref.WeakKeyDictionary)
120+ self.recomp_counter: DefaultDict[Any, int] = defaultdict(int)
121+ self.is_recomputed: DefaultDict[Any, bool] = defaultdict(bool)
122+ self.early_stop = early_stop
123+ self.metadata_fn = metadata_fn
124+ self.x_metadatas: List[Any] = []
125+ self.forward_completed = False
126+ self.ignore_saved_mismatch = False
127+ self.active_session: Optional[_SessionActivation] = None
128+ 
129+ def check_recomputed_tensors_match(self, session_id: Any) -> None:
130+ """Validate saved-tensor count and metadata after recomputation."""
131+ if self.ignore_saved_mismatch:
132+ return
133+ if len(self.weak_holders) != self.recomp_counter[session_id]:
134+ raise CheckpointError(
135+ "Hyper checkpoint saved a different number of tensors during forward and recomputation. "
136+ f"Forward saved {len(self.weak_holders)} tensors, but recomputation saved "
137+ f"{self.recomp_counter[session_id]} tensors."
138+ )
139+ 
140+ mismatches = []
141+ for index, weak_holder in enumerate(self.weak_holders):
142+ holder = weak_holder()
143+ if holder is None:
144+ continue
145+ handle = holder.handles.get(session_id)
146+ _internal_assert(handle is not None, "Missing recomputed tensor handle during metadata validation.")
147+ _internal_assert(
148+ handle in self.recomputed[session_id],
149+ "Missing recomputed tensor during metadata validation.",
150+ )
151+ recomputed_tensor = self.recomputed[session_id][handle]
152+ recomputed_metadata = self.metadata_fn(recomputed_tensor)
153+ if self.x_metadatas[index] != recomputed_metadata:
154+ mismatches.append((index, self.x_metadatas[index], recomputed_metadata))
155+ 
156+ if mismatches:
157+ details = "\n".join(
158+ f"tensor {index}: forward={forward_metadata}, recompute={recomputed_metadata}"
159+ for index, forward_metadata, recomputed_metadata in mismatches
160+ )
161+ raise CheckpointError(
162+ "Hyper checkpoint detected different tensor metadata during recomputation:\n" + details
163+ )
164+ 
165+ def clear_session(self, session_id: Any) -> None:
166+ """Release all tensors and handles associated with one session."""
167+ for weak_holder in self.weak_holders:
168+ holder = weak_holder()
169+ if holder is not None:
170+ holder.handles.pop(session_id, None)
171+ self.recomputed.pop(session_id, None)
172+ self.recomp_counter.pop(session_id, None)
173+ self.is_recomputed.pop(session_id, None)
174+ 
175+ 
176+def _bind_session_activation(frame: _CheckpointFrame, activation: _SessionActivation) -> None:
177+ """Bind one activation to a frame outside the unpack hot path."""
178+ if frame.active_session is activation:
179+ return
180+ if frame.active_session is not None:
181+ raise CheckpointError("Concurrent recompute sessions on the same checkpoint frame are not supported.")
182+ frame.active_session = activation
183+ activation.frames.add(frame)
184+ 
185+ 
186+def _register_session_frame(
187+ frame: _CheckpointFrame,
188+ session_id: Any,
189+ activation: Optional[_SessionActivation] = None,
190+) -> None:
191+ """Register a frame for cleanup and bind its current activation when present."""
192+ with _SESSION_FRAMES_LOCK:
193+ _SESSION_FRAMES[session_id].add(frame)
194+ if activation is not None:
195+ _internal_assert(activation.session_id == session_id, "Session activation key does not match its frame.")
196+ _bind_session_activation(frame, activation)
197+ 
198+ 
199+def _activate_registered_frames(activation: _SessionActivation) -> None:
200+ """Install an activation on every frame already registered for its session."""
201+ with _SESSION_FRAMES_LOCK:
202+ for frame in list(_SESSION_FRAMES.get(activation.session_id, ())):
203+ _bind_session_activation(frame, activation)
204+ 
205+ 
206+def _deactivate_session(activation: _SessionActivation) -> None:
207+ """Remove one activation from every frame bound at context entry."""
208+ with _SESSION_FRAMES_LOCK:
209+ for frame in list(activation.frames):
210+ if frame.active_session is activation:
211+ frame.active_session = None
212+ activation.frames.clear()
213+ 
214+ 
215+def _internal_assert(condition: bool, message: str) -> None:
216+ if not condition:
217+ raise CheckpointError(message)
218+ 
219+ 
220+def _noop_context_fn() -> Tuple[contextlib.AbstractContextManager, contextlib.AbstractContextManager]:
221+ return contextlib.nullcontext(), contextlib.nullcontext()
222+ 
223+ 
224+def _default_metadata_fn(tensor: Any) -> Dict[str, Any]:
225+ return {"shape": tensor.shape, "dtype": tensor.dtype, "device": tensor.device}
226+ 
227+ 
228+def _infer_device_type(*args: Any) -> str:
229+ """Return the preferred non-CPU device type found in checkpoint inputs."""
230+ device_types = []
231+ 
232+ def add_device_type(value: Any) -> None:
233+ """Record one non-CPU tensor device type."""
234+ if isinstance(value, torch.Tensor) and value.device.type != "cpu":
235+ device_types.append(value.device.type)
236+ 
237+ tree_map(add_device_type, args)
238+ device_types_set = set(device_types)
239+ if len(device_types_set) > 1:
240+ warnings.warn(
241+ "Hyper checkpoint received tensors on multiple non-CPU device types. RNG state is preserved only for "
242+ "one device type; CUDA is preferred when present.",
243+ stacklevel=3,
244+ )
245+ if not device_types:
246+ return DefaultDeviceType.get_device_type()
247+ if "cuda" in device_types_set:
248+ return "cuda"
249+ return device_types[0]
250+ 
251+ 
252+def _get_device_module(device_type: str) -> Any:
253+ if device_type == "meta":
254+ return torch.device("meta")
255+ return getattr(torch, device_type)
256+ 
257+ 
258+def _get_device_states(device_type: str, *args: Any) -> Tuple[List[int], List[Any]]:
259+ """Capture RNG states for non-CPU input devices of the requested type."""
260+ device_ids = []
261+ 
262+ def add_device_id(value: Any) -> None:
263+ """Record one non-CPU tensor device index."""
264+ if isinstance(value, torch.Tensor) and value.device.type not in {"cpu", "meta"}:
265+ device_ids.append(value.get_device())
266+ 
267+ tree_map(add_device_id, args)
268+ device_module = _get_device_module(device_type)
269+ states = []
270+ for device_id in device_ids:
271+ with device_module.device(device_id):
272+ states.append(device_module.get_rng_state())
273+ return device_ids, states
274+ 
275+ 
276+def _set_device_states(device_type: str, devices: List[int], states: List[Any]) -> None:
277+ if device_type == "meta":
278+ return
279+ device_module = _get_device_module(device_type)
280+ for device, state in zip(devices, states):
281+ with device_module.device(device):
282+ device_module.set_rng_state(state)
283+ 
284+ 
285+def _get_autocast_kwargs(device_type: str) -> Tuple[Optional[Dict[str, Any]], Dict[str, Any]]:
286+ """Return active autocast settings for the selected device and CPU."""
287+ device_kwargs = None
288+ if torch.amp.is_autocast_available(device_type):
289+ device_kwargs = {
290+ "enabled": torch.is_autocast_enabled(device_type),
291+ "dtype": torch.get_autocast_dtype(device_type),
292+ "cache_enabled": torch.is_autocast_cache_enabled(),
293+ }
294+ cpu_kwargs = {
295+ "enabled": torch.is_autocast_enabled("cpu"),
296+ "dtype": torch.get_autocast_dtype("cpu"),
297+ "cache_enabled": torch.is_autocast_cache_enabled(),
298+ }
299+ return device_kwargs, cpu_kwargs
300+ 
301+ 
302+def _create_recomputation_hooks(frame: _CheckpointFrame, session_id: Any) -> Any:
303+ """Create saved-tensor hooks that retain tensors from one recomputation."""
304+ frame_ref = weakref.ref(frame)
305+ 
306+ def pack_hook(tensor: Any) -> Any:
307+ """Store recomputed tensors in their forward holders."""
308+ tensor = tensor.detach() if tensor.requires_grad else tensor
309+ target_frame = frame_ref()
310+ _internal_assert(target_frame is not None, "Checkpoint frame was released during recomputation.")
311+ recompute_index = target_frame.recomp_counter[session_id]
312+ target_frame.recomp_counter[session_id] += 1
313+ 
314+ if recompute_index >= len(target_frame.weak_holders):
315+ if not target_frame.early_stop and not target_frame.forward_completed:
316+ target_frame.ignore_saved_mismatch = True
317+ return tensor
318+ raise CheckpointError(
319+ "Hyper checkpoint tried to save more tensors during recomputation than during forward."
320+ )
321+ 
322+ holder = target_frame.weak_holders[recompute_index]()
323+ if holder is not None:
324+ _internal_assert(
325+ holder.handles.get(session_id) is None,
326+ "A recomputed tensor handle already exists for this session.",
327+ )
328+ handle = _Handle()
329+ holder.handles[session_id] = handle
330+ target_frame.recomputed[session_id][handle] = tensor
331+ 
332+ if target_frame.early_stop and target_frame.recomp_counter[session_id] == len(target_frame.weak_holders):
333+ raise _StopRecomputationError
334+ return tensor
335+ 
336+ def unpack_hook(tensor: Any) -> Any:
337+ """Return tensors saved by operations inside the recomputation."""
338+ return tensor
339+ 
340+ return torch.autograd.graph.saved_tensors_hooks(pack_hook, unpack_hook)
341+ 
342+ 
343+# PyTorch exposes this tracing guard only as a private decorator.
344+@torch._disable_dynamo # pylint: disable=protected-access
345+def _run_fn_with_dynamo_disabled(function: Callable, *args: Any, **kwargs: Any) -> Any:
346+ """Run recomputation without tracing the saved-tensor unpack hook with Dynamo."""
347+ return function(*args, **kwargs)
348+ 
349+ 
350+def _run_recomputation(frame: _CheckpointFrame, session_id: Any) -> None:
351+ """Run and validate a frame recomputation for the given session."""
352+ if frame.is_recomputed[session_id]:
353+ return
354+ 
355+ activation = frame.active_session
356+ if activation is not None:
357+ _internal_assert(activation.session_id == session_id, "Active session key does not match recomputation key.")
358+ previous_activation = _RECOMPUTE_SESSION.get()
359+ token = None
360+ if activation is not None and previous_activation is not activation:
361+ token = _RECOMPUTE_SESSION.set(activation)
362+ try:
363+ input_context = frame.input_saver.grad_fn
364+ args = input_context.get_args(input_context.saved_tensors)
365+ try:
366+ with _create_recomputation_hooks(frame, session_id), torch.autograd.enable_grad():
367+ _run_fn_with_dynamo_disabled(frame.recompute_fn, *args)
368+ except _StopRecomputationError:
369+ pass
370+ finally:
371+ if token is not None:
372+ _RECOMPUTE_SESSION.reset(token)
373+ frame.is_recomputed[session_id] = True
374+ frame.check_recomputed_tensors_match(session_id)
375+ 
376+ 
377+def _create_checkpoint_hooks(frame: _CheckpointFrame) -> Any:
378+ """Create hooks that lazily recompute tensors saved during forward."""
379+ def pack_hook(tensor: Any) -> _Holder:
380+ """Replace a forward saved tensor with an opaque holder."""
381+ holder = _Holder()
382+ frame.weak_holders.append(weakref.ref(holder))
383+ if frame.metadata_fn is not None:
384+ with torch.no_grad():
385+ frame.x_metadatas.append(frame.metadata_fn(tensor))
386+ return holder
387+ 
388+ def unpack_hook(holder: _Holder) -> Any:
389+ """Return the corresponding tensor from lazy or prefired recomputation."""
390+ activation = frame.active_session
391+ if activation is not None:
392+ session_id = activation.session_id
393+ retain_on_unpack = activation.retain_on_unpack
394+ else:
395+ session_id = torch._C._current_graph_task_id() # pylint: disable=W0212
396+ if session_id == -1:
397+ session_id = int(uuid.uuid4())
398+ retain_on_unpack = False
399+ 
400+ _run_recomputation(frame, session_id)
401+ _internal_assert(session_id in holder.handles, "No recomputed tensor was saved for this checkpoint value.")
402+ handle = holder.handles[session_id]
403+ if handle is None:
404+ raise CheckpointError("A checkpoint tensor was unpacked more than once in the same recompute session.")
405+ _internal_assert(handle in frame.recomputed[session_id], "The recomputed tensor has already been released.")
406+ tensor = frame.recomputed[session_id][handle]
407+ if not retain_on_unpack:
408+ holder.handles[session_id] = None
409+ return tensor
410+ 
411+ return torch.autograd.graph.saved_tensors_hooks(pack_hook, unpack_hook)
412+ 
413+ 
414+def _is_compiling() -> bool:
415+ compiler = getattr(torch, "compiler", None)
416+ return bool(compiler is not None and compiler.is_compiling())
417+ 
418+ 
419+def _native_checkpoint(
420+ function: Callable,
421+ *args: Any,
422+ context_fn: Callable,
423+ preserve_rng_state: bool,
424+ determinism_check: str,
425+ debug: bool,
426+ early_stop: bool,
427+ **kwargs: Any,
428+) -> Any:
429+ """Use the public native API for compile, adapting 2.6/2.7 early-stop."""
430+ with set_checkpoint_early_stop(early_stop):
431+ return torch_checkpoint(
432+ function,
433+ *args,
434+ use_reentrant=False,
435+ context_fn=context_fn,
436+ preserve_rng_state=preserve_rng_state,
437+ determinism_check=determinism_check,
438+ debug=debug,
439+ **kwargs,
440+ )
441+ 
442+ 
443+def _checkpoint_without_reentrant_generator(
444+ function: Callable,
445+ preserve_rng_state: bool,
446+ context_fn: Callable,
447+ determinism_check: str,
448+ early_stop: bool,
449+ *args: Any,
450+ **kwargs: Any,
451+) -> Generator[None, None, None]:
452+ """Set up eager checkpoint state around the caller's forward execution."""
453+ metadata_functions = {_DEFAULT_DETERMINISM_MODE: _default_metadata_fn, "none": lambda tensor: None}
454+ if determinism_check not in metadata_functions:
455+ raise ValueError(
456+ f"determinism_check must be one of {list(metadata_functions)}, but got {determinism_check!r}."
457+ )
458+ metadata_fn = metadata_functions[determinism_check]
459+ 
460+ device_type = _infer_device_type(*args)
461+ device_module = _get_device_module(device_type)
462+ contexts = context_fn()
463+ if not isinstance(contexts, tuple) or len(contexts) != 2:
464+ raise ValueError("context_fn must return a (forward_context, recompute_context) tuple.")
465+ forward_context, recompute_context = contexts
466+ device_autocast_kwargs, cpu_autocast_kwargs = _get_autocast_kwargs(device_type)
467+ 
468+ had_device_in_forward = False
469+ forward_devices: List[int] = []
470+ forward_device_states: List[Any] = []
471+ forward_cpu_state = None
472+ if preserve_rng_state:
473+ forward_cpu_state = torch.get_rng_state()
474+ if getattr(device_module, "_initialized", False):
475+ had_device_in_forward = True
476+ forward_devices, forward_device_states = _get_device_states(device_type, *args)
477+ 
478+ def recompute_fn(*inputs: Any) -> None:
479+ """Restore execution state and rerun the checkpointed function."""
480+ function_kwargs, *function_args = inputs
481+ rng_devices = forward_devices if preserve_rng_state and had_device_in_forward else []
482+ with torch.random.fork_rng(
483+ devices=rng_devices,
484+ enabled=preserve_rng_state,
485+ device_type=device_type,
486+ ):
487+ if preserve_rng_state:
488+ torch.set_rng_state(forward_cpu_state)
489+ if had_device_in_forward:
490+ _set_device_states(device_type, forward_devices, forward_device_states)
491+ 
492+ device_autocast_context = contextlib.nullcontext()
493+ if device_autocast_kwargs is not None:
494+ device_autocast_context = torch.amp.autocast(device_type=device_type, **device_autocast_kwargs)
495+ with device_autocast_context, torch.amp.autocast("cpu", **cpu_autocast_kwargs), recompute_context:
496+ function(*function_args, **function_kwargs)
497+ 
498+ frame = _CheckpointFrame(recompute_fn, early_stop, metadata_fn)
499+ dummy = torch.empty((0,), requires_grad=True)
500+ frame.input_saver = _NoopSaveInputs.apply(dummy, kwargs, *args)
501+ 
502+ if frame.input_saver.grad_fn is None:
503+ yield
504+ return
505+ 
506+ activation = _RECOMPUTE_SESSION.get()
507+ if activation is not None:
508+ raise CheckpointError("Nested checkpoint is not supported during scheduled recomputation.")
509+ 
510+ collector = _RECOMPUTE_COLLECTOR.get()
511+ if collector is not None:
512+ collector.append(frame)
513+ try:
514+ with _create_checkpoint_hooks(frame), forward_context:
515+ yield
516+ frame.forward_completed = True
517+ 
518+ if getattr(device_module, "_initialized", False) and preserve_rng_state and not had_device_in_forward:
519+ raise RuntimeError(
520+ "The device state was initialized inside a Hyper checkpoint forward, so its initial RNG state "
521+ "could not be preserved. Initialize the device before entering checkpoint."
522+ )
523+ except BaseException:
524+ if collector is not None and frame in collector:
525+ collector.remove(frame)
526+ raise
527+ 
528+ 
529+def checkpoint(
530+ function: Callable,
531+ *args: Any,
532+ use_reentrant: bool = False,
533+ context_fn: Callable = _noop_context_fn,
534+ preserve_rng_state: bool = True,
535+ determinism_check: str = _DEFAULT_DETERMINISM_MODE,
536+ debug: bool = False,
537+ early_stop: bool = True,
538+ **kwargs: Any,
539+) -> Any:
540+ """Run Hyper's non-reentrant checkpoint implementation.
541+ 
542+ Eager execution always uses this implementation. Compile execution falls
543+ back to PyTorch's public non-reentrant checkpoint API.
544+ """
545+ if use_reentrant is not False:
546+ raise ValueError("Hyper checkpoint only supports use_reentrant=False.")
547+ if not isinstance(early_stop, bool):
548+ raise ValueError(f"early_stop must be bool, but got {type(early_stop).__name__}.")
549+ if not isinstance(preserve_rng_state, bool):
550+ raise ValueError(
551+ f"preserve_rng_state must be bool, but got {type(preserve_rng_state).__name__}."
552+ )
553+ if not callable(context_fn):
554+ raise ValueError("context_fn must be callable.")
555+ if _is_compiling():
556+ return _native_checkpoint(
557+ function,
558+ *args,
559+ context_fn=context_fn,
560+ preserve_rng_state=preserve_rng_state,
561+ determinism_check=determinism_check,
562+ debug=debug,
563+ early_stop=early_stop,
564+ **kwargs,
565+ )
566+ if debug:
567+ raise ValueError("debug=True is not supported by Hyper eager checkpoint yet.")
568+ 
569+ generator = _checkpoint_without_reentrant_generator(
570+ function,
571+ preserve_rng_state,
572+ context_fn,
573+ determinism_check,
574+ early_stop,
575+ *args,
576+ **kwargs,
577+ )
578+ next(generator)
579+ try:
580+ result = function(*args, **kwargs)
581+ except BaseException:
582+ generator.close()
583+ raise
584+ try:
585+ next(generator)
586+ except StopIteration:
587+ return result
588+ generator.close()
589+ raise CheckpointError("The internal checkpoint generator yielded more than once.")
590+ 
591+ 
592+@contextlib.contextmanager
593+def recompute_handle_collector_ctx() -> Iterator[List[Any]]:
594+ """Collect opaque checkpoint handles created in this context."""
595+ handles = []
596+ token = _RECOMPUTE_COLLECTOR.set(handles)
597+ try:
598+ yield handles
599+ finally:
600+ _RECOMPUTE_COLLECTOR.reset(token)
601+ 
602+ 
603+def recompute_handle(handle: Any, session_id: Any) -> None:
604+ """Run one collected checkpoint recomputation ahead of backward."""
605+ if not isinstance(handle, _CheckpointFrame):
606+ raise ValueError("handle must be produced by recompute_handle_collector_ctx().")
607+ _validate_session_id(session_id)
608+ activation = _RECOMPUTE_SESSION.get()
609+ if (
610+ activation is not None
611+ and activation.session_id == session_id
612+ and activation.retain_on_unpack
613+ ):
614+ _register_session_frame(handle, session_id, activation)
615+ _run_recomputation(handle, session_id)
616+ return
617+ if activation is not None:
618+ raise CheckpointError("recompute_handle cannot enter another active recompute session.")
619+ 
620+ _register_session_frame(handle, session_id)
621+ with recompute_session_ctx(session_id=session_id, retain_on_unpack=True):
622+ _run_recomputation(handle, session_id)
623+ 
624+ 
625+def _validate_session_id(session_id: Any) -> None:
626+ if session_id is None:
627+ raise ValueError("session_id must not be None.")
628+ try:
629+ hash(session_id)
630+ except TypeError as error:
631+ raise ValueError("session_id must be hashable.") from error
632+ 
633+ 
634+@contextlib.contextmanager
635+def recompute_session_ctx(session_id: Any, retain_on_unpack: bool = False) -> Iterator[Any]:
636+ """Select the key and retention policy used by checkpoint unpack hooks."""
637+ _validate_session_id(session_id)
638+ if not isinstance(retain_on_unpack, bool):
639+ raise ValueError(f"retain_on_unpack must be bool, but got {type(retain_on_unpack).__name__}.")
640+ if _RECOMPUTE_SESSION.get() is not None:
641+ raise CheckpointError("Nested recompute session contexts are not supported.")
642+ activation = _SessionActivation(session_id, retain_on_unpack)
643+ token = _RECOMPUTE_SESSION.set(activation)
644+ try:
645+ _activate_registered_frames(activation)
646+ yield session_id
647+ finally:
648+ try:
649+ _deactivate_session(activation)
650+ finally:
651+ _RECOMPUTE_SESSION.reset(token)
652+ 
653+ 
654+def clear_recompute_session(session_id: Any) -> None:
655+ """Release retained recomputation data for a session; repeated calls are safe."""
656+ _validate_session_id(session_id)
657+ with _SESSION_FRAMES_LOCK:
658+ registered_frames = _SESSION_FRAMES.pop(session_id, None)
659+ frames = list(registered_frames) if registered_frames is not None else []
660+ for frame in frames:
661+ frame.clear_session(session_id)
@@ -1468,7 +1468,33 @@ class TorchPlatform(Platform):
1468 1468 
1469 @property1469 @property
1470 def checkpoint(self):1470 def checkpoint(self):
1471- return torch.utils.checkpoint.checkpoint1471+ # pylint: disable=C0415
1472+ from hyper_parallel.platform.torch.activation_checkpoint.checkpoint import checkpoint
1473+ return checkpoint
1474+ 
1475+ @staticmethod
1476+ def recompute_handle_collector_ctx():
1477+ # pylint: disable=C0415
1478+ from hyper_parallel.platform.torch.activation_checkpoint.checkpoint import recompute_handle_collector_ctx
1479+ return recompute_handle_collector_ctx()
1480+ 
1481+ @staticmethod
1482+ def recompute_handle(handle, session_id):
1483+ # pylint: disable=C0415
1484+ from hyper_parallel.platform.torch.activation_checkpoint.checkpoint import recompute_handle
1485+ return recompute_handle(handle, session_id)
1486+ 
1487+ @staticmethod
1488+ def recompute_session_ctx(session_id, retain_on_unpack=False):
1489+ # pylint: disable=C0415
1490+ from hyper_parallel.platform.torch.activation_checkpoint.checkpoint import recompute_session_ctx
1491+ return recompute_session_ctx(session_id=session_id, retain_on_unpack=retain_on_unpack)
1492+ 
1493+ @staticmethod
1494+ def clear_recompute_session(session_id):
1495+ # pylint: disable=C0415
1496+ from hyper_parallel.platform.torch.activation_checkpoint.checkpoint import clear_recompute_session
1497+ return clear_recompute_session(session_id)
1472 1498 
1473 @staticmethod1499 @staticmethod
1474 def checkpoint_wrapper(module, **checkpoint_kwargs):1500 def checkpoint_wrapper(module, **checkpoint_kwargs):
@@ -17,17 +17,223 @@ import copy
17import multiprocessing as mp17import multiprocessing as mp
18import queue18import queue
19import traceback19import traceback
20+from typing import Any
20 21 
21import pytest22import pytest
22import torch23import torch
24+from torch.utils.checkpoint import DefaultDeviceType
25+from torch.utils.checkpoint import checkpoint as torch_checkpoint
23 26 
24-from hyper_parallel.core.activation_checkpoint import CheckpointPolicy, SwapManager, checkpoint_wrapper, swap_wrapper27+from hyper_parallel.core.activation_checkpoint import (
28+ CheckpointPolicy,
29+ SwapManager,
30+ checkpoint,
31+ checkpoint_wrapper,
32+ swap_wrapper,
33+)
34+from hyper_parallel.platform import get_platform
35+from hyper_parallel.platform.torch.activation_checkpoint import checkpoint as hyper_checkpoint
25from tests.torch.common_net import SimpleTransformer36from tests.torch.common_net import SimpleTransformer
26from tests.torch.activation_checkpoint.utils import prepare_data, seed_memory_time_context, set_seed, train_one_mode37from tests.torch.activation_checkpoint.utils import prepare_data, seed_memory_time_context, set_seed, train_one_mode
27 38 
28 39 
29MEMORY_COMPARISON_MODES = ("none", "recompute", "save", "swap", "group_swap")40MEMORY_COMPARISON_MODES = ("none", "recompute", "save", "swap", "group_swap")
30MEMORY_COMPARISON_VOCAB_SIZE = 819241MEMORY_COMPARISON_VOCAB_SIZE = 8192
42+platform = get_platform()
43+ 
44+ 
45+def _prefire_recompute(handles: list, session_id: Any) -> None:
46+ """Run all collected checkpoint frames in one retained session."""
47+ with platform.recompute_session_ctx(session_id=session_id, retain_on_unpack=True):
48+ for handle in handles:
49+ platform.recompute_handle(handle, session_id)
50+ 
51+ 
52+def _scheduled_dx_dw(
53+ output: torch.Tensor,
54+ input_tensor: torch.Tensor,
55+ weights: tuple,
56+ session_id: Any,
57+) -> tuple:
58+ """Compute separate input and weight gradients from one prefired session."""
59+ with platform.recompute_session_ctx(session_id=session_id, retain_on_unpack=True):
60+ input_grad = torch.autograd.grad(output.sum(), input_tensor, retain_graph=True)[0]
61+ with platform.recompute_session_ctx(session_id=session_id, retain_on_unpack=False):
62+ weight_grads = torch.autograd.grad(output.sum(), weights)
63+ return input_grad, weight_grads
64+ 
65+ 
66+def _run_closure_only_npu_rng_case(checkpoint_fn):
67+ """Return checkpoint output and gradient when the NPU tensor is captured only by closure."""
68+ captured_tensor = torch.ones(4096, device="npu", requires_grad=True)
69+ cpu_argument = torch.zeros(1)
70+ torch.npu.manual_seed(2026)
71+ 
72+ def function(unused_argument: torch.Tensor) -> torch.Tensor:
73+ """Run NPU dropout while the device tensor is captured only by closure."""
74+ del unused_argument
75+ return torch.nn.functional.dropout(captured_tensor, p=0.5, training=True)
76+ 
77+ output = checkpoint_fn(function, cpu_argument)
78+ output.sum().backward()
79+ return output.detach(), captured_tensor.grad.detach()
80+ 
81+ 
82+def test_native_npu_rng_for_closure_only_tensor_is_not_restored():
83+ """Native checkpoint does not discover an NPU tensor that exists only in a closure."""
84+ previous_device_type = DefaultDeviceType.get_device_type()
85+ DefaultDeviceType.set_device_type("npu")
86+ try:
87+ native_output, native_grad = _run_closure_only_npu_rng_case(
88+ lambda function, argument: torch_checkpoint(function, argument, use_reentrant=False)
89+ )
90+ finally:
91+ DefaultDeviceType.set_device_type(previous_device_type)
92+ 
93+ assert not torch.equal(native_grad, native_output)
94+ 
95+ 
96+def test_hyper_npu_rng_for_closure_only_tensor_matches_native():
97+ """Hyper should match native RNG behavior when the NPU tensor exists only in a closure."""
98+ previous_device_type = DefaultDeviceType.get_device_type()
99+ DefaultDeviceType.set_device_type("npu")
100+ try:
101+ native_output, native_grad = _run_closure_only_npu_rng_case(
102+ lambda function, argument: torch_checkpoint(function, argument, use_reentrant=False)
103+ )
104+ hyper_output, hyper_grad = _run_closure_only_npu_rng_case(hyper_checkpoint)
105+ finally:
106+ DefaultDeviceType.set_device_type(previous_device_type)
107+ 
108+ assert not torch.equal(native_grad, native_output)
109+ assert torch.equal(hyper_output, native_output)
110+ assert torch.equal(hyper_grad, native_grad)
111+ 
112+ 
113+def test_scheduled_recompute_supports_dx_dw_split() -> None:
114+ """A prefired recomputation should be shared by separate dx and dw autograd calls."""
115+ set_seed(2026)
116+ input_tensor = torch.randn(8, 16, device="npu", requires_grad=True)
117+ weight = torch.randn(16, 16, device="npu", requires_grad=True)
118+ reference_input = input_tensor.detach().clone().requires_grad_()
119+ reference_weight = weight.detach().clone().requires_grad_()
120+ 
121+ reference_output = torch.nn.functional.gelu(reference_input @ reference_weight)
122+ expected_dx, expected_dw = torch.autograd.grad(
123+ reference_output.sum(),
124+ (reference_input, reference_weight),
125+ )
126+ 
127+ checkpoint_calls = 0
128+ 
129+ def checkpointed_function(current_input: torch.Tensor, current_weight: torch.Tensor) -> torch.Tensor:
130+ """Count checkpoint executions while applying a weight-dependent function."""
131+ nonlocal checkpoint_calls
132+ checkpoint_calls += 1
133+ return torch.nn.functional.gelu(current_input @ current_weight)
134+ 
135+ with platform.recompute_handle_collector_ctx() as handles:
136+ output = checkpoint(checkpointed_function, input_tensor, weight)
137+ 
138+ assert len(handles) == 1
139+ session_id = ("dxdw_split_e2e", id(output))
140+ try:
141+ _prefire_recompute(handles, session_id)
142+ actual_dx, actual_dws = _scheduled_dx_dw(output, input_tensor, (weight,), session_id)
143+ finally:
144+ platform.clear_recompute_session(session_id)
145+ 
146+ assert checkpoint_calls == 2
147+ torch.testing.assert_close(output, reference_output)
148+ torch.testing.assert_close(actual_dx, expected_dx)
149+ torch.testing.assert_close(actual_dws[0], expected_dw)
150+ 
151+ 
152+def test_scheduled_recompute_npu_preserves_rng_state() -> None:
153+ """NPU dropout should replay its forward mask without advancing caller RNG state."""
154+ set_seed(2030)
155+ input_tensor = torch.randn(16, 32, device="npu", requires_grad=True)
156+ weight = torch.randn(32, 32, device="npu", requires_grad=True)
157+ reference_input = input_tensor.detach().clone().requires_grad_()
158+ reference_weight = weight.detach().clone().requires_grad_()
159+ 
160+ set_seed(88)
161+ reference_output = torch.nn.functional.dropout(
162+ torch.nn.functional.gelu(reference_input @ reference_weight),
163+ p=0.4,
164+ training=True,
165+ )
166+ expected_dx, expected_dw = torch.autograd.grad(
167+ reference_output.sum(),
168+ (reference_input, reference_weight),
169+ )
170+ 
171+ set_seed(88)
172+ with platform.recompute_handle_collector_ctx() as handles:
173+ output = checkpoint(
174+ lambda current_input, current_weight: torch.nn.functional.dropout(
175+ torch.nn.functional.gelu(current_input @ current_weight),
176+ p=0.4,
177+ training=True,
178+ ),
179+ input_tensor,
180+ weight,
181+ preserve_rng_state=True,
182+ )
183+ 
184+ session_id = ("npu_rng", id(output))
185+ rng_state_before_prefire = torch.npu.get_rng_state()
186+ try:
187+ _prefire_recompute(handles, session_id)
188+ rng_state_after_prefire = torch.npu.get_rng_state()
189+ actual_dx, actual_dws = _scheduled_dx_dw(output, input_tensor, (weight,), session_id)
190+ finally:
191+ platform.clear_recompute_session(session_id)
192+ 
193+ assert torch.equal(rng_state_after_prefire, rng_state_before_prefire)
194+ torch.testing.assert_close(output, reference_output)
195+ torch.testing.assert_close(actual_dx, expected_dx)
196+ torch.testing.assert_close(actual_dws[0], expected_dw)
197+ 
198+ 
199+def test_scheduled_recompute_npu_restores_autocast() -> None:
200+ """A prefire outside autocast should restore the NPU forward autocast settings."""
201+ set_seed(2031)
202+ input_tensor = torch.randn(16, 32, device="npu", requires_grad=True)
203+ weight = torch.randn(32, 32, device="npu", requires_grad=True)
204+ reference_input = input_tensor.detach().clone().requires_grad_()
205+ reference_weight = weight.detach().clone().requires_grad_()
206+ 
207+ with torch.autocast(device_type="npu", dtype=torch.bfloat16):
208+ reference_output = torch.nn.functional.gelu(reference_input @ reference_weight)
209+ expected_dx, expected_dw = torch.autograd.grad(
210+ reference_output.float().sum(),
211+ (reference_input, reference_weight),
212+ )
213+ execution_dtypes = []
214+ 
215+ def checkpointed_function(current_input: torch.Tensor, current_weight: torch.Tensor) -> torch.Tensor:
216+ """Record the matmul dtype used in forward and scheduled replay."""
217+ result = current_input @ current_weight
218+ execution_dtypes.append(result.dtype)
219+ return torch.nn.functional.gelu(result)
220+ 
221+ with torch.autocast(device_type="npu", dtype=torch.bfloat16):
222+ with platform.recompute_handle_collector_ctx() as handles:
223+ output = checkpoint(checkpointed_function, input_tensor, weight)
224+ 
225+ session_id = ("npu_autocast", id(output))
226+ try:
227+ _prefire_recompute(handles, session_id)
228+ actual_dx, actual_dws = _scheduled_dx_dw(output.float(), input_tensor, (weight,), session_id)
229+ finally:
230+ platform.clear_recompute_session(session_id)
231+ 
232+ assert execution_dtypes == [torch.bfloat16, torch.bfloat16]
233+ assert output.dtype == torch.bfloat16
234+ torch.testing.assert_close(output, reference_output)
235+ torch.testing.assert_close(actual_dx, expected_dx, atol=5e-3, rtol=5e-3)
236+ torch.testing.assert_close(actual_dws[0], expected_dw, atol=5e-3, rtol=5e-3)
31 237 
32 238 
33def apply_recompute(model, mode):239def apply_recompute(model, mode):
@@ -0,0 +1,671 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""NPU system-test cases for Hyper's Torch non-reentrant checkpoint implementation."""
16+# pylint: disable=missing-public-docstring,missing-public-type-hints,wrong-import-position,protected-access
17+import contextlib
18+import importlib
19+import inspect
20+import os
21+import unittest
22+from typing import Callable, List
23+from unittest.mock import patch
24+ 
25+import torch
26+ 
27+os.environ["HYPER_PARALLEL_PLATFORM"] = "torch"
28+ 
29+from hyper_parallel.platform.torch.activation_checkpoint.checkpoint import (
30+ CheckpointError,
31+ checkpoint,
32+ clear_recompute_session,
33+ recompute_handle,
34+ recompute_handle_collector_ctx,
35+ recompute_session_ctx,
36+)
37+from hyper_parallel.core.activation_checkpoint import CheckpointPolicy, checkpoint as core_checkpoint
38+from tests.torch.activation_checkpoint.utils import set_seed
39+ 
40+checkpoint_module = importlib.import_module("hyper_parallel.platform.torch.activation_checkpoint.checkpoint")
41+ 
42+_NPU_DEVICE = "npu"
43+ 
44+ 
45+def _randn(*size: int, **kwargs) -> torch.Tensor:
46+ """Create a random tensor explicitly on NPU."""
47+ return torch.randn(*size, device=_NPU_DEVICE, **kwargs)
48+ 
49+ 
50+def _ones(*size: int, **kwargs) -> torch.Tensor:
51+ """Create an all-ones tensor explicitly on NPU."""
52+ return torch.ones(*size, device=_NPU_DEVICE, **kwargs)
53+ 
54+ 
55+def _make_tail_recording_function(
56+ function: Callable[..., torch.Tensor], tail_calls: List[None]
57+) -> Callable[..., torch.Tensor]:
58+ """Wrap a test function and record each completed invocation."""
59+ def wrapped(*args: torch.Tensor) -> torch.Tensor:
60+ """Run the wrapped function and append one tail-call marker."""
61+ result = function(*args)
62+ tail_calls.append(None)
63+ return result
64+ 
65+ return wrapped
66+ 
67+ 
68+class TestCheckpoint(unittest.TestCase):
69+ """Validate eager checkpoint semantics and input checks."""
70+ 
71+ def test_backward_matches_non_checkpointed_function(self):
72+ """Checkpoint gradients should match regular autograd for inputs and weights."""
73+ input_tensor = _randn(4, 4, requires_grad=True)
74+ weight = _randn(4, 4, requires_grad=True)
75+ baseline_input = input_tensor.detach().clone().requires_grad_()
76+ baseline_weight = weight.detach().clone().requires_grad_()
77+ 
78+ baseline = torch.sin(baseline_input * baseline_weight)
79+ baseline.sum().backward()
80+ 
81+ output = checkpoint(lambda x, w: torch.sin(x * w), input_tensor, weight)
82+ output.sum().backward()
83+ 
84+ self.assertTrue(torch.allclose(input_tensor.grad, baseline_input.grad))
85+ self.assertTrue(torch.allclose(weight.grad, baseline_weight.grad))
86+ 
87+ def test_function_keyword_arguments_are_forwarded(self):
88+ """Non-checkpoint control kwargs should be passed to the wrapped function."""
89+ input_tensor = _randn(4, requires_grad=True)
90+ 
91+ output = checkpoint(lambda x, scale: x.sin() * scale, input_tensor, scale=3.0)
92+ output.sum().backward()
93+ 
94+ self.assertTrue(torch.allclose(output, input_tensor.sin() * 3.0))
95+ 
96+ def test_early_stop_controls_tail_execution(self):
97+ """early_stop should stop before side effects after the last saved tensor."""
98+ for early_stop, expected_tail_calls in ((True, 1), (False, 2)):
99+ with self.subTest(early_stop=early_stop):
100+ tail_calls = []
101+ 
102+ function = _make_tail_recording_function(
103+ lambda input_tensor: input_tensor.sin().cos(), tail_calls
104+ )
105+ 
106+ input_tensor = _randn(4, requires_grad=True)
107+ checkpoint(function, input_tensor, early_stop=early_stop).sum().backward()
108+ 
109+ self.assertEqual(len(tail_calls), expected_tail_calls)
110+ 
111+ def test_native_early_stop_context_does_not_override_keyword(self):
112+ """The eager implementation should use only Hyper's explicit keyword."""
113+ tail_calls = []
114+ 
115+ def function(input_tensor):
116+ output = input_tensor.sin().cos()
117+ tail_calls.append(None)
118+ return output
119+ 
120+ input_tensor = _randn(4, requires_grad=True)
121+ with torch.utils.checkpoint.set_checkpoint_early_stop(False):
122+ checkpoint(function, input_tensor, early_stop=True).sum().backward()
123+ 
124+ self.assertEqual(len(tail_calls), 1)
125+ 
126+ def test_preserve_rng_state_matches_regular_autograd(self):
127+ """Random masks should be restored when preserve_rng_state is enabled."""
128+ baseline_input = _ones(32, requires_grad=True)
129+ checkpoint_input = baseline_input.detach().clone().requires_grad_()
130+ 
131+ set_seed(7)
132+ baseline_output = torch.nn.functional.dropout(baseline_input, p=0.5, training=True)
133+ baseline_output.sum().backward()
134+ 
135+ set_seed(7)
136+ output = checkpoint(
137+ lambda x: torch.nn.functional.dropout(x, p=0.5, training=True),
138+ checkpoint_input,
139+ preserve_rng_state=True,
140+ )
141+ output.sum().backward()
142+ 
143+ self.assertTrue(torch.equal(output, baseline_output))
144+ self.assertTrue(torch.equal(checkpoint_input.grad, baseline_input.grad))
145+ 
146+ def test_forward_and_recompute_contexts_are_used(self):
147+ """The two contexts should surround their corresponding executions."""
148+ events = []
149+ 
150+ @contextlib.contextmanager
151+ def record(name):
152+ events.append(f"enter:{name}")
153+ try:
154+ yield
155+ finally:
156+ events.append(f"exit:{name}")
157+ 
158+ def context_fn():
159+ return record("forward"), record("recompute")
160+ 
161+ input_tensor = _randn(4, requires_grad=True)
162+ checkpoint(lambda x: x.sin(), input_tensor, context_fn=context_fn).sum().backward()
163+ 
164+ self.assertEqual(
165+ events,
166+ ["enter:forward", "exit:forward", "enter:recompute", "exit:recompute"],
167+ )
168+ 
169+ def test_grad_inside_forward_matches_regular_autograd(self):
170+ """Checkpoint should support an unpack triggered by grad inside forward."""
171+ input_tensor = _randn(4, requires_grad=True)
172+ baseline_input = input_tensor.detach().clone().requires_grad_()
173+ 
174+ def function(value):
175+ intermediate = value.sin()
176+ inner_grad = torch.autograd.grad(intermediate.sum(), value, create_graph=True)[0]
177+ return intermediate * inner_grad
178+ 
179+ baseline_output = function(baseline_input)
180+ baseline_output.sum().backward()
181+ output = checkpoint(function, input_tensor)
182+ output.sum().backward()
183+ 
184+ self.assertTrue(torch.allclose(output, baseline_output))
185+ self.assertTrue(torch.allclose(input_tensor.grad, baseline_input.grad))
186+ 
187+ def test_no_grad_execution_does_not_create_handle(self):
188+ """A checkpoint outside autograd should behave as a direct function call."""
189+ with recompute_handle_collector_ctx() as handles:
190+ with torch.no_grad():
191+ output = checkpoint(lambda x: x.sin(), _ones(4))
192+ 
193+ self.assertTrue(torch.equal(output, _ones(4).sin()))
194+ self.assertEqual(handles, [])
195+ 
196+ def test_failed_forward_does_not_leave_handle(self):
197+ """A failed checkpoint invocation should be removed from its collector."""
198+ def function(input_tensor):
199+ del input_tensor
200+ raise RuntimeError("forward failed")
201+ 
202+ with recompute_handle_collector_ctx() as handles:
203+ with self.assertRaisesRegex(RuntimeError, "forward failed"):
204+ checkpoint(function, _ones(1, requires_grad=True))
205+ 
206+ self.assertEqual(handles, [])
207+ 
208+ def test_invalid_checkpoint_options_raise(self):
209+ """Unsupported modes and malformed options should fail at the API boundary."""
210+ input_tensor = _ones(1, requires_grad=True)
211+ with self.assertRaisesRegex(ValueError, "use_reentrant=False"):
212+ checkpoint(lambda x: x, input_tensor, use_reentrant=True)
213+ with self.assertRaisesRegex(ValueError, "early_stop must be bool"):
214+ checkpoint(lambda x: x, input_tensor, early_stop=1)
215+ with self.assertRaisesRegex(ValueError, "preserve_rng_state must be bool"):
216+ checkpoint(lambda x: x, input_tensor, preserve_rng_state=1)
217+ with self.assertRaisesRegex(ValueError, "determinism_check"):
218+ checkpoint(lambda x: x, input_tensor, determinism_check="invalid")
219+ with self.assertRaisesRegex(ValueError, "debug=True"):
220+ checkpoint(lambda x: x, input_tensor, debug=True)
221+ 
222+ def test_compile_uses_native_checkpoint(self):
223+ """Compile state should fall back to the public native checkpoint API."""
224+ with patch.object(checkpoint_module, "_is_compiling", return_value=True), patch.object(
225+ checkpoint_module, "_native_checkpoint", return_value="compiled-result"
226+ ) as mock_native_checkpoint:
227+ result = checkpoint(lambda x: x, 1, early_stop=False)
228+ 
229+ self.assertEqual(result, "compiled-result")
230+ self.assertFalse(mock_native_checkpoint.call_args.kwargs["early_stop"])
231+ 
232+ 
233+class TestScheduledRecomputation(unittest.TestCase):
234+ """Validate early scheduling and dx/dw separated backward behavior."""
235+ 
236+ def test_separate_dx_dw_without_session_recomputes_per_graph_task(self):
237+ """Native GraphTask keys should avoid unpack loss across dx and dw calls."""
238+ calls = []
239+ input_tensor = _randn(4, 4, requires_grad=True)
240+ weight = _randn(4, 4, requires_grad=True)
241+ 
242+ def function(x, w):
243+ calls.append(None)
244+ return torch.nn.functional.gelu(x @ w)
245+ 
246+ output = checkpoint(function, input_tensor, weight)
247+ torch.autograd.grad(output.sum(), input_tensor, retain_graph=True)
248+ torch.autograd.grad(output.sum(), weight)
249+ 
250+ self.assertEqual(len(calls), 3)
251+ 
252+ def test_prefired_session_is_shared_by_dx_and_dw(self):
253+ """A retained session should reuse one prefired recomputation for dx and dw."""
254+ calls = []
255+ input_tensor = _randn(4, 4, requires_grad=True)
256+ weight = _randn(4, 4, requires_grad=True)
257+ baseline_input = input_tensor.detach().clone().requires_grad_()
258+ baseline_weight = weight.detach().clone().requires_grad_()
259+ 
260+ baseline = torch.nn.functional.gelu(baseline_input @ baseline_weight)
261+ expected_dx = torch.autograd.grad(baseline.sum(), baseline_input, retain_graph=True)[0]
262+ expected_dw = torch.autograd.grad(baseline.sum(), baseline_weight)[0]
263+ 
264+ def function(x, w):
265+ calls.append(None)
266+ return torch.nn.functional.gelu(x @ w)
267+ 
268+ with recompute_handle_collector_ctx() as handles:
269+ output = checkpoint(function, input_tensor, weight)
270+ 
271+ session_id = ("micro-batch", 0)
272+ try:
273+ with recompute_session_ctx(session_id, retain_on_unpack=True):
274+ recompute_handle(handles[0], session_id)
275+ actual_dx = torch.autograd.grad(output.sum(), input_tensor, retain_graph=True)[0]
276+ actual_dw = torch.autograd.grad(output.sum(), weight)[0]
277+ finally:
278+ clear_recompute_session(session_id)
279+ 
280+ self.assertEqual(len(handles), 1)
281+ self.assertEqual(len(calls), 2)
282+ self.assertTrue(torch.allclose(actual_dx, expected_dx))
283+ self.assertTrue(torch.allclose(actual_dw, expected_dw))
284+ 
285+ def test_one_session_serves_multiple_checkpoint_frames(self):
286+ """One session should prefire and serve multiple sequential checkpoint frames."""
287+ input_tensor = _randn(4, 4, requires_grad=True)
288+ first_weight = _randn(4, 6, requires_grad=True)
289+ second_weight = _randn(6, 3, requires_grad=True)
290+ reference_input = input_tensor.detach().clone().requires_grad_()
291+ reference_first_weight = first_weight.detach().clone().requires_grad_()
292+ reference_second_weight = second_weight.detach().clone().requires_grad_()
293+ reference_hidden = torch.nn.functional.gelu(reference_input @ reference_first_weight)
294+ reference_output = torch.nn.functional.silu(reference_hidden @ reference_second_weight)
295+ expected_grads = torch.autograd.grad(
296+ reference_output.sum(),
297+ (reference_input, reference_first_weight, reference_second_weight),
298+ )
299+ calls = [0, 0]
300+ 
301+ def first_block(current_input, current_weight):
302+ calls[0] += 1
303+ return torch.nn.functional.gelu(current_input @ current_weight)
304+ 
305+ def second_block(current_input, current_weight):
306+ calls[1] += 1
307+ return torch.nn.functional.silu(current_input @ current_weight)
308+ 
309+ with recompute_handle_collector_ctx() as handles:
310+ hidden = checkpoint(first_block, input_tensor, first_weight)
311+ output = checkpoint(second_block, hidden, second_weight)
312+ 
313+ session_id = ("multiple-frames", id(output))
314+ try:
315+ with recompute_session_ctx(session_id, retain_on_unpack=True):
316+ for handle in handles:
317+ recompute_handle(handle, session_id)
318+ with recompute_session_ctx(session_id, retain_on_unpack=True):
319+ actual_dx = torch.autograd.grad(output.sum(), input_tensor, retain_graph=True)[0]
320+ with recompute_session_ctx(session_id, retain_on_unpack=False):
321+ actual_dws = torch.autograd.grad(output.sum(), (first_weight, second_weight))
322+ finally:
323+ clear_recompute_session(session_id)
324+ 
325+ self.assertEqual(len(handles), 2)
326+ self.assertEqual(calls, [2, 2])
327+ self.assertTrue(torch.allclose(actual_dx, expected_grads[0]))
328+ self.assertTrue(torch.allclose(actual_dws[0], expected_grads[1]))
329+ self.assertTrue(torch.allclose(actual_dws[1], expected_grads[2]))
330+ 
331+ def test_repeated_sessions_keep_iterations_isolated(self):
332+ """Repeated iterations should not reuse frames or tensors from earlier sessions."""
333+ weight = _randn(4, 4, requires_grad=True)
334+ calls = []
335+ 
336+ def function(current_input, current_weight):
337+ calls.append(None)
338+ return torch.nn.functional.gelu(current_input @ current_weight)
339+ 
340+ for step in range(3):
341+ input_tensor = _randn(4, 4, requires_grad=True)
342+ reference_input = input_tensor.detach().clone().requires_grad_()
343+ reference_weight = weight.detach().clone().requires_grad_()
344+ reference_output = torch.nn.functional.gelu(reference_input @ reference_weight)
345+ expected_dx, expected_dw = torch.autograd.grad(
346+ reference_output.sum(),
347+ (reference_input, reference_weight),
348+ )
349+ with recompute_handle_collector_ctx() as handles:
350+ output = checkpoint(function, input_tensor, weight)
351+ 
352+ session_id = ("iteration", step, id(output))
353+ try:
354+ recompute_handle(handles[0], session_id)
355+ with recompute_session_ctx(session_id, retain_on_unpack=True):
356+ actual_dx = torch.autograd.grad(output.sum(), input_tensor, retain_graph=True)[0]
357+ with recompute_session_ctx(session_id, retain_on_unpack=False):
358+ actual_dw = torch.autograd.grad(output.sum(), weight)[0]
359+ finally:
360+ clear_recompute_session(session_id)
361+ clear_recompute_session(session_id)
362+ 
363+ self.assertTrue(torch.allclose(actual_dx, expected_dx))
364+ self.assertTrue(torch.allclose(actual_dw, expected_dw))
365+ self.assertEqual(len(calls), (step + 1) * 2)
366+ 
367+ def test_partial_and_failed_sessions_can_be_cleared_and_reused(self):
368+ """Early-stop, partial consumption, and failed prefire should leave reusable frames."""
369+ for early_stop, expected_tail_calls in ((True, 1), (False, 3)):
370+ with self.subTest(early_stop=early_stop):
371+ input_tensor = _randn(4, 4, requires_grad=True)
372+ weight = _randn(4, 4, requires_grad=True)
373+ reference_input = input_tensor.detach().clone().requires_grad_()
374+ reference_weight = weight.detach().clone().requires_grad_()
375+ reference_output = torch.nn.functional.gelu(reference_input @ reference_weight)
376+ expected_dx, expected_dw = torch.autograd.grad(
377+ reference_output.sum(),
378+ (reference_input, reference_weight),
379+ )
380+ tail_calls = []
381+ 
382+ function = _make_tail_recording_function(
383+ lambda current_input, current_weight: torch.nn.functional.gelu(current_input @ current_weight),
384+ tail_calls,
385+ )
386+ 
387+ with recompute_handle_collector_ctx() as handles:
388+ output = checkpoint(function, input_tensor, weight, early_stop=early_stop)
389+ 
390+ session_id = ("partial", early_stop, id(output))
391+ try:
392+ recompute_handle(handles[0], session_id)
393+ with recompute_session_ctx(session_id, retain_on_unpack=True):
394+ actual_dx = torch.autograd.grad(output.sum(), input_tensor, retain_graph=True)[0]
395+ finally:
396+ clear_recompute_session(session_id)
397+ 
398+ actual_dw = torch.autograd.grad(output.sum(), weight)[0]
399+ self.assertEqual(len(tail_calls), expected_tail_calls)
400+ self.assertTrue(torch.allclose(actual_dx, expected_dx))
401+ self.assertTrue(torch.allclose(actual_dw, expected_dw))
402+ 
403+ input_tensor = _randn(4, requires_grad=True)
404+ should_fail = True
405+ calls = []
406+ 
407+ def failing_function(value):
408+ calls.append(None)
409+ result = value.sin().cos()
410+ if should_fail and len(calls) > 1:
411+ raise RuntimeError("expected recomputation failure")
412+ return result
413+ 
414+ with recompute_handle_collector_ctx() as handles:
415+ output = checkpoint(failing_function, input_tensor, early_stop=False)
416+ 
417+ failed_session_id = ("failed", id(output))
418+ try:
419+ with self.assertRaisesRegex(RuntimeError, "expected recomputation failure"):
420+ recompute_handle(handles[0], failed_session_id)
421+ finally:
422+ clear_recompute_session(failed_session_id)
423+ 
424+ should_fail = False
425+ recovered_session_id = ("recovered", id(output))
426+ try:
427+ recompute_handle(handles[0], recovered_session_id)
428+ with recompute_session_ctx(recovered_session_id, retain_on_unpack=False):
429+ actual_grad = torch.autograd.grad(output.sum(), input_tensor)[0]
430+ finally:
431+ clear_recompute_session(recovered_session_id)
432+ 
433+ reference_input = input_tensor.detach().clone().requires_grad_()
434+ expected_grad = torch.autograd.grad(reference_input.sin().cos().sum(), reference_input)[0]
435+ self.assertEqual(len(calls), 3)
436+ self.assertTrue(torch.allclose(actual_grad, expected_grad))
437+ 
438+ def test_prefired_session_survives_contextvar_loss(self):
439+ """A worker without the caller ContextVar should still find prefired tensors."""
440+ calls = []
441+ input_tensor = _randn(4, 4, requires_grad=True)
442+ weight = _randn(4, 4, requires_grad=True)
443+ 
444+ def function(x, w):
445+ calls.append(None)
446+ return torch.nn.functional.gelu(x @ w)
447+ 
448+ with recompute_handle_collector_ctx() as handles:
449+ output = checkpoint(function, input_tensor, weight)
450+ 
451+ session_id = "worker-context-loss"
452+ try:
453+ recompute_handle(handles[0], session_id)
454+ with recompute_session_ctx(session_id, retain_on_unpack=True):
455+ token = checkpoint_module._RECOMPUTE_SESSION.set(None)
456+ try:
457+ torch.autograd.grad(output.sum(), input_tensor, retain_graph=True)
458+ finally:
459+ checkpoint_module._RECOMPUTE_SESSION.reset(token)
460+ with recompute_session_ctx(session_id, retain_on_unpack=False):
461+ token = checkpoint_module._RECOMPUTE_SESSION.set(None)
462+ try:
463+ torch.autograd.grad(output.sum(), weight)
464+ finally:
465+ checkpoint_module._RECOMPUTE_SESSION.reset(token)
466+ finally:
467+ clear_recompute_session(session_id)
468+ 
469+ self.assertEqual(len(calls), 2)
470+ 
471+ def test_scheduled_nested_checkpoint_raises(self):
472+ """Scheduled recomputation should reject nested checkpoint regions explicitly."""
473+ input_tensor = _randn(4, requires_grad=True)
474+ 
475+ def inner(value):
476+ return value.sin()
477+ 
478+ def outer(value):
479+ return checkpoint(inner, value).cos()
480+ 
481+ with recompute_handle_collector_ctx() as handles:
482+ checkpoint(outer, input_tensor)
483+ 
484+ session_id = "nested-prefire"
485+ try:
486+ with self.assertRaisesRegex(CheckpointError, "Nested checkpoint is not supported"):
487+ recompute_handle(handles[0], session_id)
488+ finally:
489+ clear_recompute_session(session_id)
490+ 
491+ self.assertEqual(len(handles), 2)
492+ 
493+ def test_nested_checkpoint_without_session_still_works(self):
494+ """Ordinary GraphTask-based recomputation should preserve native nested behavior."""
495+ input_tensor = _randn(4, requires_grad=True)
496+ baseline_input = input_tensor.detach().clone().requires_grad_()
497+ 
498+ def inner(value):
499+ return value.sin()
500+ 
501+ def outer(value):
502+ return checkpoint(inner, value).cos()
503+ 
504+ expected = baseline_input.sin().cos()
505+ expected_grad = torch.autograd.grad(expected.sum(), baseline_input)[0]
506+ output = checkpoint(outer, input_tensor)
507+ actual_grad = torch.autograd.grad(output.sum(), input_tensor)[0]
508+ 
509+ self.assertTrue(torch.allclose(output, expected))
510+ self.assertTrue(torch.allclose(actual_grad, expected_grad))
511+ 
512+ def test_prefired_session_consumes_sac_cache_only_once(self):
513+ """SAC replay should happen during prefire, not again for dx and dw consumers."""
514+ calls = []
515+ input_tensor = _randn(4, 4, requires_grad=True)
516+ weight = _randn(4, 4, requires_grad=True)
517+ 
518+ def function(x, w):
519+ calls.append(None)
520+ return torch.nn.functional.gelu(x @ w)
521+ 
522+ def policy_fn(ctx, op, *args, **kwargs):
523+ del ctx, op, args, kwargs
524+ return CheckpointPolicy.MUST_SAVE
525+ 
526+ with recompute_handle_collector_ctx() as handles:
527+ output = core_checkpoint(function, input_tensor, weight, policy_fn=policy_fn)
528+ 
529+ session_id = "sac-split-backward"
530+ try:
531+ recompute_handle(handles[0], session_id)
532+ with recompute_session_ctx(session_id, retain_on_unpack=True):
533+ torch.autograd.grad(output.sum(), input_tensor, retain_graph=True)
534+ torch.autograd.grad(output.sum(), weight)
535+ finally:
536+ clear_recompute_session(session_id)
537+ 
538+ self.assertEqual(len(calls), 2)
539+ 
540+ def test_clear_session_is_idempotent(self):
541+ """A retained session can be cleared repeatedly in cleanup paths."""
542+ input_tensor = _randn(4, requires_grad=True)
543+ with recompute_handle_collector_ctx() as handles:
544+ checkpoint(lambda x: x.sin(), input_tensor)
545+ 
546+ session_id = "clear-twice"
547+ recompute_handle(handles[0], session_id)
548+ clear_recompute_session(session_id)
549+ clear_recompute_session(session_id)
550+ 
551+ def test_invalid_handle_and_session_options_raise(self):
552+ """Scheduling APIs should reject handles and session options from outside Hyper."""
553+ with self.assertRaisesRegex(ValueError, "handle must be produced"):
554+ recompute_handle(object(), "session")
555+ with self.assertRaisesRegex(ValueError, "session_id must be hashable"):
556+ with recompute_session_ctx([], retain_on_unpack=False):
557+ pass
558+ with self.assertRaisesRegex(ValueError, "retain_on_unpack must be bool"):
559+ with recompute_session_ctx("session", retain_on_unpack=1):
560+ pass
561+ 
562+ def test_session_id_is_required_and_must_not_be_none(self):
563+ """Session contexts should require callers to propagate one explicit stable id."""
564+ session_parameter = inspect.signature(recompute_session_ctx).parameters["session_id"]
565+ self.assertIs(session_parameter.default, inspect.Parameter.empty)
566+ with self.assertRaisesRegex(ValueError, "session_id must not be None"):
567+ with recompute_session_ctx(None):
568+ pass
569+ 
570+ def test_unpack_outside_backward_uses_temporary_graph_key(self):
571+ """An eager saved-tensor access outside backward should recompute successfully."""
572+ input_tensor = _randn(4, requires_grad=True)
573+ output = checkpoint(lambda value: value.sin(), input_tensor)
574+ 
575+ saved_input = output.grad_fn._saved_self # pylint: disable=W0212
576+ 
577+ self.assertTrue(torch.equal(saved_input, input_tensor))
578+ 
579+ def test_default_device_type_is_used_without_device_tensor_arguments(self):
580+ """Device-less arguments should honor Torch's stable checkpoint default."""
581+ with patch.object(checkpoint_module.DefaultDeviceType, "get_device_type", return_value="npu"):
582+ device_type = checkpoint_module._infer_device_type(torch.ones(1, device="cpu"))
583+ 
584+ self.assertEqual(device_type, "npu")
585+ 
586+ def test_incomplete_forward_ignores_extra_save_only_without_early_stop(self):
587+ """Only a full recomputation may run ahead of an incomplete original forward."""
588+ input_tensor = _randn(4, requires_grad=True)
589+ early_stop_frame = checkpoint_module._CheckpointFrame(lambda: None, True, None)
590+ with self.assertRaises(CheckpointError):
591+ with checkpoint_module._create_recomputation_hooks(early_stop_frame, "early-stop"):
592+ input_tensor.sin()
593+ 
594+ full_recompute_frame = checkpoint_module._CheckpointFrame(lambda: None, False, None)
595+ with checkpoint_module._create_recomputation_hooks(full_recompute_frame, "full-recompute"):
596+ input_tensor.sin()
597+ self.assertTrue(full_recompute_frame.ignore_saved_mismatch)
598+ 
599+ def test_recompute_metadata_mismatch_raises(self):
600+ """Default determinism checks should detect changed recompute tensor metadata."""
601+ recomputing = False
602+ 
603+ @contextlib.contextmanager
604+ def recompute_context():
605+ nonlocal recomputing
606+ recomputing = True
607+ try:
608+ yield
609+ finally:
610+ recomputing = False
611+ 
612+ def function(input_tensor):
613+ value = input_tensor.half() if recomputing else input_tensor
614+ return value.sin()
615+ 
616+ input_tensor = _randn(4, requires_grad=True)
617+ output = checkpoint(
618+ function,
619+ input_tensor,
620+ context_fn=lambda: (contextlib.nullcontext(), recompute_context()),
621+ )
622+ 
623+ with self.assertRaises(CheckpointError):
624+ output.sum().backward()
625+ 
626+ def test_none_determinism_check_allows_metadata_change(self):
627+ """The none mode should keep count checks while skipping tensor metadata checks."""
628+ recomputing = False
629+ 
630+ @contextlib.contextmanager
631+ def recompute_context():
632+ nonlocal recomputing
633+ recomputing = True
634+ try:
635+ yield
636+ finally:
637+ recomputing = False
638+ 
639+ def function(input_tensor):
640+ value = input_tensor.half() if recomputing else input_tensor
641+ return value.sin()
642+ 
643+ input_tensor = _randn(4, requires_grad=True)
644+ output = checkpoint(
645+ function,
646+ input_tensor,
647+ context_fn=lambda: (contextlib.nullcontext(), recompute_context()),
648+ determinism_check="none",
649+ )
650+ 
651+ output.sum().backward()
652+ 
653+ 
654+def run_checkpoint_cases() -> None:
655+ """Run all checkpoint cases on NPU."""
656+ loader = unittest.TestLoader()
657+ suite = unittest.TestSuite(
658+ (
659+ loader.loadTestsFromTestCase(TestCheckpoint),
660+ loader.loadTestsFromTestCase(TestScheduledRecomputation),
661+ )
662+ )
663+ result = unittest.TextTestRunner(verbosity=2).run(suite)
664+ if not result.wasSuccessful():
665+ raise AssertionError(
666+ f"Checkpoint NPU suite failed with {len(result.failures)} failures and {len(result.errors)} errors."
667+ )
668+ 
669+ 
670+if __name__ == "__main__":
671+ run_checkpoint_cases()
@@ -15,6 +15,8 @@
15"""test activation checkpoint"""15"""test activation checkpoint"""
16from tests.common.mark_utils import arg_mark16from tests.common.mark_utils import arg_mark
17from tests.common.parallel_case import parallel_run, TorchCase17from tests.common.parallel_case import parallel_run, TorchCase
18+from tests.torch.activation_checkpoint import activation_checkpoint as activation_checkpoint_cases
19+from tests.torch.activation_checkpoint import checkpoint_cases
18 20 
19ACTIVATION_CHECKPOINT = "activation_checkpoint.py"21ACTIVATION_CHECKPOINT = "activation_checkpoint.py"
20 22 
@@ -36,3 +38,39 @@ def test_ac_memory_group():
36 TorchCase(ACTIVATION_CHECKPOINT, "test_wrapper_overlap_detection_cases", 12406, 1),38 TorchCase(ACTIVATION_CHECKPOINT, "test_wrapper_overlap_detection_cases", 12406, 1),
37 TorchCase(ACTIVATION_CHECKPOINT, "test_wrapper_non_overlapping_allowed_cases", 12407, 1)39 TorchCase(ACTIVATION_CHECKPOINT, "test_wrapper_non_overlapping_allowed_cases", 12407, 1)
38 ])40 ])
41+ 
42+ 
43+@arg_mark(plat_marks=["platform_ascend910b"], level_mark="level0", card_mark="onecard", essential_mark="essential")
44+def test_native_npu_rng_for_closure_only_tensor_is_not_restored():
45+ """Verify native checkpoint does not restore closure-only NPU RNG state."""
46+ activation_checkpoint_cases.test_native_npu_rng_for_closure_only_tensor_is_not_restored()
47+ 
48+ 
49+@arg_mark(plat_marks=["platform_ascend910b"], level_mark="level0", card_mark="onecard", essential_mark="essential")
50+def test_hyper_npu_rng_for_closure_only_tensor_matches_native():
51+ """Verify Hyper matches native closure-only NPU RNG behavior."""
52+ activation_checkpoint_cases.test_hyper_npu_rng_for_closure_only_tensor_matches_native()
53+ 
54+ 
55+@arg_mark(plat_marks=["platform_ascend910b"], level_mark="level0", card_mark="onecard", essential_mark="essential")
56+def test_scheduled_recompute_supports_dx_dw_split() -> None:
57+ """Verify one prefired recomputation serves separate dx and dw autograd calls."""
58+ activation_checkpoint_cases.test_scheduled_recompute_supports_dx_dw_split()
59+ 
60+ 
61+@arg_mark(plat_marks=["platform_ascend910b"], level_mark="level0", card_mark="onecard", essential_mark="essential")
62+def test_scheduled_recompute_npu_preserves_rng_state() -> None:
63+ """Verify scheduled NPU recomputation preserves random state."""
64+ activation_checkpoint_cases.test_scheduled_recompute_npu_preserves_rng_state()
65+ 
66+ 
67+@arg_mark(plat_marks=["platform_ascend910b"], level_mark="level0", card_mark="onecard", essential_mark="essential")
68+def test_scheduled_recompute_npu_restores_autocast() -> None:
69+ """Verify scheduled NPU recomputation restores autocast settings."""
70+ activation_checkpoint_cases.test_scheduled_recompute_npu_restores_autocast()
71+ 
72+ 
73+@arg_mark(plat_marks=["platform_ascend910b"], level_mark="level0", card_mark="onecard", essential_mark="essential")
74+def test_checkpoint_npu_semantics() -> None:
75+ """Run eager checkpoint semantics and scheduling coverage on NPU."""
76+ checkpoint_cases.run_checkpoint_cases()
@@ -13,6 +13,9 @@
13# limitations under the License.13# limitations under the License.
14# ============================================================================14# ============================================================================
15"""Unit tests for activation checkpoint module."""15"""Unit tests for activation checkpoint module."""
16+# The backend selector must be set before importing platform aliases. The
17+# local imports and patched platform fixture are intentional test setup.
18+# pylint: disable=wrong-import-position,import-outside-toplevel,unused-argument
16import contextlib19import contextlib
17import os20import os
18import unittest21import unittest
@@ -161,6 +164,21 @@ class TestCheckpointFunction(unittest.TestCase):
161 call_args = mock_plat.checkpoint.call_args[0]164 call_args = mock_plat.checkpoint.call_args[0]
162 self.assertIn(3, call_args)165 self.assertIn(3, call_args)
163 166 
167+ def test_checkpoint_forwards_early_stop_keyword(self, mock_plat):
168+ """Test checkpoint forwards the explicit early_stop control keyword."""
169+ mock_plat.checkpoint.return_value = "result"
170+ 
171+ result = checkpoint(lambda value: value, 3, **{"early_stop": False})
172+ 
173+ self.assertEqual(result, "result")
174+ self.assertFalse(mock_plat.checkpoint.call_args.kwargs["early_stop"])
175+ 
176+ def test_checkpoint_rejects_non_boolean_early_stop(self, mock_plat):
177+ """Test checkpoint rejects ambiguous early_stop values."""
178+ with self.assertRaisesRegex(ValueError, "early_stop must be bool"):
179+ checkpoint(lambda value: value, 3, early_stop=1)
180+ mock_plat.checkpoint.assert_not_called()
181+ 
164 def test_checkpoint_composes_recompute_state_and_user_contexts(self, mock_plat):182 def test_checkpoint_composes_recompute_state_and_user_contexts(self, mock_plat):
165 """Unified recompute state should surround user checkpoint contexts."""183 """Unified recompute state should surround user checkpoint contexts."""
166 events = []184 events = []
@@ -13,6 +13,9 @@
13# limitations under the License.13# limitations under the License.
14# ============================================================================14# ============================================================================
15"""Unit tests for PyTorch activation checkpoint wrapper."""15"""Unit tests for PyTorch activation checkpoint wrapper."""
16+# These tests intentionally select the Torch backend before importing aliases,
17+# patch the platform object, and inspect wrapper state for API coverage.
18+# pylint: disable=wrong-import-position,protected-access,unused-argument,cyclic-import
16import os19import os
17import unittest20import unittest
18from unittest.mock import patch21from unittest.mock import patch
@@ -133,6 +136,17 @@ class TestCheckpointWrapper(unittest.TestCase):
133 mock_plat.create_selective_checkpoint_contexts.assert_called_once_with(136 mock_plat.create_selective_checkpoint_contexts.assert_called_once_with(
134 policy, group_swap=False)137 policy, group_swap=False)
135 138 
139+ def test_do_checkpoint_passes_early_stop(self, mock_plat):
140+ """Test wrapper forwards its early_stop configuration to core checkpoint."""
141+ mock_plat.checkpoint.return_value = "result"
142+ mod = _BaseWrapperModule()
143+ wrapper = CheckpointWrapper(mod, early_stop=False)
144+ 
145+ result = wrapper.forward(torch.randn(2, 4))
146+ 
147+ self.assertEqual(result, "result")
148+ self.assertFalse(mock_plat.checkpoint.call_args.kwargs["early_stop"])
149+ 
136 150 
137class TestCkptWrapper(unittest.TestCase):151class TestCkptWrapper(unittest.TestCase):
138 """Unit tests for PyTorch ckpt_wrapper() factory function."""152 """Unit tests for PyTorch ckpt_wrapper() factory function."""