已合并
feat: add scheduled torch activation recomputation #1102
DavidFFFan创建于 8月1日
feat: add scheduled torch activation recomputation #1102
已合并
共 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( | |||
| 890 | checkpoint_wrapper(module, **checkpoint_kwargs) -> CheckpointWrapper | 898 | checkpoint_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 | ||
| 55 | output = checkpoint(model.layer, x, context_fn=context_fn) | 55 | output = 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. 函数式 swap | 74 | ### 2. 函数式 swap |
| 59 | 75 | ||
| 60 | ```python | 76 | ```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.nullcontext | 155 | 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, **kwargs | 158 | + 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 | 442 | ||
| 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 the | 445 | # 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 | 614 | ||
| 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_name | 617 | hook_name = ctx.hook_name |
| 618 | coordinator = ctx.coordinator | 618 | 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 | 659 | ||
| 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 | 697 | ||
| 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 | 1863 | ||
| 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=C0415 | 1867 | # pylint: disable=C0415 |
| 1866 | from mindspore.common.recompute import _recompute_session_ctx | 1868 | 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 | ||
| 16 | import os | 19 | import os |
| 17 | from datetime import timedelta | 20 | from datetime import timedelta |
| 18 | from enum import auto, Enum | 21 | from 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 this | 1554 | + session_id: Required stable session key. Recompute caches are keyed |
| 1552 | - instead of the transient autodiff engine id, so a re-run fired | 1555 | + 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 recomputed | 1558 | retain_on_unpack (bool): When ``True``, unpack returns recomputed |
| 1555 | tensors without popping them, so a later backward can consume | 1559 | 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""" |
| 16 | from .checkpoint_wrapper import CheckpointWrapper, ckpt_wrapper | 16 | from .checkpoint_wrapper import CheckpointWrapper, ckpt_wrapper |
| 17 | from .activation_swap import swap_wrapper, swap_tensor_wrapper | 17 | from .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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | +# 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 | + | ||
| 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 | + | ||
| 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 | 1469 | ||
| 1470 | def checkpoint(self): | 1470 | def checkpoint(self): |
| 1471 | - return torch.utils.checkpoint.checkpoint | 1471 | + # pylint: disable=C0415 |
| 1472 | + from hyper_parallel.platform.torch.activation_checkpoint.checkpoint import checkpoint | ||
| 1473 | + return checkpoint | ||
| 1474 | + | ||
| 1475 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | 1499 | ||
| 1474 | def checkpoint_wrapper(module, **checkpoint_kwargs): | 1500 | def checkpoint_wrapper(module, **checkpoint_kwargs): |
| @@ -17,17 +17,223 @@ import copy | |||
| 17 | import multiprocessing as mp | 17 | import multiprocessing as mp |
| 18 | import queue | 18 | import queue |
| 19 | import traceback | 19 | import traceback |
| 20 | +from typing import Any | ||
| 20 | 21 | ||
| 21 | import pytest | 22 | import pytest |
| 22 | import torch | 23 | import 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_wrapper | 27 | +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 | ||
| 25 | from tests.torch.common_net import SimpleTransformer | 36 | from tests.torch.common_net import SimpleTransformer |
| 26 | from tests.torch.activation_checkpoint.utils import prepare_data, seed_memory_time_context, set_seed, train_one_mode | 37 | from tests.torch.activation_checkpoint.utils import prepare_data, seed_memory_time_context, set_seed, train_one_mode |
| 27 | 38 | ||
| 28 | 39 | ||
| 29 | MEMORY_COMPARISON_MODES = ("none", "recompute", "save", "swap", "group_swap") | 40 | MEMORY_COMPARISON_MODES = ("none", "recompute", "save", "swap", "group_swap") |
| 30 | MEMORY_COMPARISON_VOCAB_SIZE = 8192 | 41 | MEMORY_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 | ||
| 33 | def apply_recompute(model, mode): | 239 | def 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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""" |
| 16 | from tests.common.mark_utils import arg_mark | 16 | from tests.common.mark_utils import arg_mark |
| 17 | from tests.common.parallel_case import parallel_run, TorchCase | 17 | from 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 | ||
| 19 | ACTIVATION_CHECKPOINT = "activation_checkpoint.py" | 21 | ACTIVATION_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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | ||
| 16 | import contextlib | 19 | import contextlib |
| 17 | import os | 20 | import os |
| 18 | import unittest | 21 | import 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 | ||
| 16 | import os | 19 | import os |
| 17 | import unittest | 20 | import unittest |
| 18 | from unittest.mock import patch | 21 | from 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 | ||
| 137 | class TestCkptWrapper(unittest.TestCase): | 151 | class TestCkptWrapper(unittest.TestCase): |
| 138 | """Unit tests for PyTorch ckpt_wrapper() factory function.""" | 152 | """Unit tests for PyTorch ckpt_wrapper() factory function.""" |