Spec-Zone.ru › PyTorch 2.14

torch.utils.checkpoint

Создано: 16 июня 2025 г. | Последнее обновление: 27 июля 2026 г.

Примечание

Контрольные точки реализованы путём повторного выполнения сегмента прямого прохода для каждого сегмента с контрольной точкой во время обратного распространения. Это может привести к тому, что сохраняемые состояния, например состояние RNG, будут продвинуты дальше, чем без контрольных точек. По умолчанию контрольные точки содержат логику управления состоянием RNG, чтобы проходы с контрольными точками, использующие RNG (например, через dropout), выдавали детерминированный результат по сравнению с проходами без контрольных точек. Логика сохранения и восстановления состояний RNG может привести к умеренному снижению производительности в зависимости от времени выполнения операций с контрольными точками. Если детерминированный результат по сравнению с проходами без контрольных точек не требуется, передайте preserve_rng_state=False в checkpoint или checkpoint_sequential, чтобы не сохранять и не восстанавливать состояние RNG при каждой контрольной точке.

Логика сохранения сохраняет и восстанавливает состояние RNG для CPU и ещё одного типа устройств (тип устройства определяется по аргументам Tensor, за исключением тензоров CPU, с помощью _infer_device_type) в run_fn. Если устройств несколько, состояние будет сохранено только для устройств одного типа, а остальные устройства будут проигнорированы. Следовательно, если в каких-либо функциях с контрольными точками используются случайные значения, это может привести к неверным градиентам. (Обратите внимание: если среди обнаруженных устройств есть устройства CUDA, они будут иметь приоритет; в противном случае будет выбрано первое обнаруженное устройство.) Если тензоров CPU нет, состояние устройства по умолчанию (значение по умолчанию — cuda; его можно изменить на другое устройство с помощью DefaultDeviceType) будет сохранено и восстановлено. Однако логика не может предвидеть, будет ли пользователь перемещать Tensor на новое устройство непосредственно внутри run_fn. Поэтому, если вы переместите Tensor на новое устройство («новое» означает не входящее в набор [текущее устройство + устройства аргументов Tensor]) внутри run_fn, детерминированный результат по сравнению с проходами без контрольных точек никогда не гарантируется.

checkpoint также можно вызвать без функции, чтобы получить декоратор. Это позволяет отделить конфигурацию контрольной точки от именованных аргументов, передаваемых функции с контрольной точкой:

checkpointed_fn = checkpoint(use_reentrant=False, preserve_rng_state=False)(fn)
out = checkpointed_fn(*args, **kwargs)
torch.utils.checkpoint.checkpoint(function: Callable[[...], _T], *args: Any, use_reentrant: bool | None = None, preserve_rng_state: bool = True, context_fn: Callable[[], Tuple[ContextManager, ContextManager]] = noop_context_fn, determinism_check: str = _DEFAULT_DETERMINISM_MODE, debug: bool = False, early_stop: bool = True, respect_saved_tensors_hooks: bool | None = None, **kwargs: Any) → _T [исходный код]
torch.utils.checkpoint.checkpoint(function:None=None, *, use_reentrant:bool|None=None, preserve_rng_state:bool=True, context_fn:Callable[[],Tuple[ContextManager,ContextManager]]=noop_context_fn, determinism_check:str=_DEFAULT_DETERMINISM_MODE, debug:bool=False, early_stop:bool=True, respect_saved_tensors_hooks:bool|None=None) → Callable[[Callable[[_P],_T]],Callable[[_P],_T]]

Создаёт контрольную точку для модели или её части.

Контрольная точка активаций — это метод, позволяющий обменивать вычислительные ресурсы на память. По умолчанию тензоры, вычисленные во время прямого прохода, сохраняются до тех пор, пока не будут использованы при вычислении градиентов в обратном проходе. Чтобы уменьшить расход памяти, тензоры, созданные в переданном function, не сохраняются до обратного прохода. Вместо этого сохраняются переданные тензоры в args, а несохранённые тензоры при необходимости вычисления градиентов повторно вычисляются путём повторного вызова function в обратном проходе. Контрольные точки активаций можно применять к любой части модели — иногда это описывают как создание «контрольной точки» для этой части модели.

В настоящее время доступны две реализации контрольных точек; выбор определяется параметром use_reentrant. Рекомендуется использовать use_reentrant=False. Различия между реализациями обсуждаются в примечании ниже.

Предупреждение

Если вызов function во время обратного прохода отличается от вызова в прямом проходе, например из-за глобальной переменной, версия с контрольной точкой может работать иначе. Это может привести к возникновению ошибки или незаметно привести к неверным градиентам.

Предупреждение

Параметр use_reentrant следует передавать явно. В версии 2.9 будет вызываться исключение, если use_reentrant не передан. Если вы используете вариант use_reentrant=True, ознакомьтесь с приведённым ниже примечанием, где указаны важные особенности и возможные ограничения.

Примечание

Реентерабельный вариант контрольной точки (use_reentrant=True) и нереентерабельный вариант контрольной точки (use_reentrant=False) различаются следующим образом:

  • Нереентерабельная контрольная точка прекращает повторные вычисления, как только будут повторно вычислены все необходимые промежуточные активации. Эта возможность включена по умолчанию, но её можно отключить с помощью set_checkpoint_early_stop(). Реентерабельная контрольная точка всегда повторно вычисляет function полностью во время обратного прохода.
  • Реентерабельный вариант не записывает граф autograd во время прямого прохода, поскольку прямой проход выполняется внутри torch.no_grad(). Нереентерабельный вариант записывает граф autograd, что позволяет выполнять обратное распространение по графу внутри областей с контрольными точками.
  • Для обратного прохода реентерабельная контрольная точка поддерживает только API torch.autograd.backward() без аргумента inputs, тогда как нереентерабельный вариант поддерживает все способы выполнения обратного прохода.
  • Для реентерабельного варианта как минимум один вход и один выход должны иметь requires_grad=True. Если это условие не выполнено, для части модели с контрольной точкой градиенты вычисляться не будут. Для нереентерабельного варианта это требование отсутствует.
  • Реентерабельный вариант не учитывает тензоры во вложенных структурах (например, пользовательских объектах, списках, словарях и т. д.) как участвующие в autograd, тогда как нереентерабельный вариант учитывает.
  • Реентерабельная контрольная точка не поддерживает области с контрольными точками, содержащие тензоры, отсоединённые от вычислительного графа; нереентерабельный вариант их поддерживает. Для реентерабельного варианта, если сегмент с контрольной точкой содержит тензоры, отсоединённые с помощью detach() или с помощью torch.no_grad(), обратный проход вызовет ошибку. Это связано с тем, что checkpoint задаёт для всех выходов требование вычисления градиентов, что вызывает проблемы, если в модели определён тензор, для которого градиенты не вычисляются. Чтобы избежать этого, отсоединяйте тензоры за пределами функции checkpoint.
Параметры:
  • function – описывает, что следует выполнить в прямом проходе модели или её части. Функция также должна уметь обрабатывать входные данные, переданные в виде кортежа. Например, в LSTM, если пользователь передаёт (activation, hidden), function должна правильно использовать первый вход как activation, а второй — как hidden. Если аргумент не указан, checkpoint возвращает декоратор, которым можно обернуть функцию перед передачей аргументов пользователя.
  • args – кортеж, содержащий входные данные для function
Именованные аргументы:
  • preserve_rng_state (bool, необязательный) – не сохранять и не восстанавливать состояние RNG при каждой контрольной точке. Обратите внимание: при использовании torch.compile этот флаг не действует, и состояние RNG всегда сохраняется. Значение по умолчанию: True
  • use_reentrant (bool) – указывает, следует ли использовать вариант контрольной точки активаций, которому требуется реентерабельный autograd. Этот параметр следует передавать явно. В версии 2.9 будет вызываться исключение, если use_reentrant не передан. Если use_reentrant=False, checkpoint будет использовать реализацию, которой не требуется реентерабельный autograd. Благодаря этому checkpoint поддерживает дополнительные возможности, например корректную работу с torch.autograd.grad и передачу именованных аргументов функции с контрольной точкой.
  • context_fn (Callable, необязательный) – вызываемый объект, возвращающий кортеж из двух менеджеров контекста. Функция и её повторное вычисление будут выполняться соответственно в контекстах первого и второго менеджеров. Этот аргумент поддерживается только при условии use_reentrant=False.
  • determinism_check (str, необязательный) – строка, задающая проверку детерминированности. По умолчанию установлено значение "default", при котором сравниваются формы, типы данных и устройства повторно вычисленных тензоров с сохранёнными тензорами. Чтобы отключить эту проверку, укажите "none". В настоящее время поддерживаются только эти два значения. Если вы хотите предложить дополнительные проверки детерминированности, создайте задачу. Этот аргумент поддерживается только при условии use_reentrant=False; если use_reentrant=True, проверка детерминированности всегда отключена.
  • debug (bool, необязательный) – если True, сообщения об ошибках также будут содержать трассировку операторов, выполненных во время исходного прямого вычисления и повторного вычисления. Этот аргумент поддерживается только при условии use_reentrant=False.
  • early_stop (bool, необязательный) – если True, нереентерабельная контрольная точка прекращает повторное вычисление сразу после вычисления всех необходимых тензоров. Этот аргумент игнорируется, если use_reentrant=True. Глобальное значение можно переопределить с помощью менеджера контекста set_checkpoint_early_stop(). Значение по умолчанию: True.
  • respect_saved_tensors_hooks (bool, необязательный) – следует ли передавать тензоры, сохраняемые выборочной контрольной точкой активаций (SAC), через окружающие пользовательские хуки torch.autograd.graph.saved_tensors_hooks() (например, torch.autograd.graph.save_on_cpu()). Такие тензоры хранятся вне графа autograd, поэтому исторически эти хуки их не видели; в одном из будущих выпусков контрольные точки будут учитывать хуки по умолчанию. До тех пор значение по умолчанию (None) сохраняет прежнее поведение и выводит FutureWarning, если действует пользовательский хук; передайте True, чтобы включить эту возможность, или False, чтобы сохранить прежнее поведение без предупреждения. Не влияет на работу без SAC (обычная контрольная точка повторно вычисляет данные, а не сохраняет их) или при условии use_reentrant=True. Этот аргумент поддерживается только при условии use_reentrant=False.
Возвращает:

Результат выполнения function над *args или декоратор, если function не указан.

Пример

>>> # Direct call: checkpoint the function immediately.
>>> out = torch.utils.checkpoint.checkpoint(
...     fn, *args, use_reentrant=False, **kwargs
... )
>>> # Decorator/curried form: configure once, then call.
>>> checkpointed_fn = torch.utils.checkpoint.checkpoint(
...     use_reentrant=False,
...     preserve_rng_state=False,
... )(fn)
>>> out = checkpointed_fn(*args, **kwargs)
>>> # The same decorator can be applied directly to a function or method.
>>> @torch.utils.checkpoint.checkpoint(use_reentrant=False)
... def fn(*args, **kwargs):
...     ...
>>> out = fn(*args, **kwargs)
torch.utils.checkpoint.checkpoint_sequential(functions, segments, input, use_reentrant=None, **kwargs) [исходный код]

Создаёт контрольную точку для последовательной модели, чтобы сэкономить память.

Последовательные модели выполняют список модулей/функций по порядку. Поэтому такую модель можно разделить на несколько сегментов и создать контрольную точку для каждого из них. Промежуточные активации не сохраняются для всех сегментов, кроме последнего. Входные данные каждого сегмента с контрольной точкой будут сохранены, чтобы повторно выполнить сегмент в обратном проходе.

Предупреждение

Параметр use_reentrant следует передавать явно. В версии 2.9 будет вызываться исключение, если use_reentrant не передан. Если вы используете вариант use_reentrant=True` variant, please see :func:`~torch.utils.checkpoint.checkpoint` for the important considerations and limitations of this variant. It is recommended that you use ``use_reentrant=False.

Параметры:
  • functions – объект torch.nn.Sequential или список модулей либо функций (составляющих модель), которые нужно выполнить последовательно.
  • segments – количество частей, на которые следует разделить модель
  • input – входной тензор для functions
  • preserve_rng_state (bool, необязательный) – не сохранять и не восстанавливать состояние RNG при каждой контрольной точке. Значение по умолчанию: True
  • use_reentrant (bool) – указывает, следует ли использовать вариант контрольной точки активаций, которому требуется реентерабельный autograd. Этот параметр следует передавать явно. В версии 2.5 будет вызываться исключение, если use_reentrant не передан. Если use_reentrant=False, checkpoint будет использовать реализацию, которой не требуется реентерабельный autograd. Благодаря этому checkpoint поддерживает дополнительные возможности, например корректную работу с torch.autograd.grad и передачу именованных аргументов функции с контрольной точкой.
Возвращает:

Результат последовательного выполнения functions над *inputs

Пример

>>> model = nn.Sequential(...)
>>> input_var = checkpoint_sequential(model, chunks, input_var)
torch.utils.checkpoint.set_checkpoint_debug_enabled(enabled) [исходный код]

Менеджер контекста, задающий, следует ли контрольной точке выводить дополнительные отладочные сведения во время выполнения. Дополнительную информацию о флаге debug для checkpoint() см. в документации. Обратите внимание: при задании этого параметра менеджер контекста переопределяет значение debug, переданное контрольной точке. Чтобы использовать локальную настройку, передайте None в этот контекст.

Параметры:

enabled (bool) – следует ли контрольной точке выводить отладочную информацию. Значение по умолчанию — «None».

class torch.utils.checkpoint.CheckpointPolicy(value) [исходный код]

Перечисление для задания политики контрольных точек при обратном распространении.

Поддерживаются следующие политики:

  • {MUST,PREFER}_SAVE: результат операции сохраняется во время прямого прохода и не будет повторно вычисляться во время обратного прохода
  • {MUST,PREFER}_RECOMPUTE: результат операции не сохраняется во время прямого прохода и будет повторно вычислен во время обратного прохода
  • {MUST,PREFER}_CPU_OFFLOAD: результат операции сохраняется во время прямого прохода, перемещается в CPU и повторно загружается в GPU во время обратного прохода

Используйте MUST_* вместо PREFER_*, чтобы указать, что другие подсистемы, например torch.compile, не должны переопределять эту политику.

Примечание

Функция политики, всегда возвращающая PREFER_RECOMPUTE, эквивалентна стандартному созданию контрольных точек.

Функция политики, возвращающая PREFER_SAVE для каждой операции, НЕ эквивалентна работе без контрольных точек. Такая политика сохранит дополнительные тензоры, в том числе те, которые фактически не нужны для вычисления градиентов.

class torch.utils.checkpoint.SelectiveCheckpointContext(*, is_recompute, op_output=None) [исходный код]

Контекст, передаваемый функции политики при выборочном создании контрольных точек.

Этот класс используется для передачи соответствующих метаданных функции политики при выборочном создании контрольных точек.

Функция политики вызывается только во время прямого прохода. При повторном вычислении кэшированные значения извлекаются по индексу, поэтому is_recompute считается устаревшим и всегда имеет значение False.

Пример

>>>
>>> def policy_fn(ctx, op, *args, **kwargs):
>>>    print(ctx.op_output)
>>>
>>> context_fn = functools.partial(create_selective_checkpoint_contexts, policy_fn)
>>>
>>> out = torch.utils.checkpoint.checkpoint(
>>>     fn, x, y,
>>>     use_reentrant=False,
>>>     context_fn=context_fn,
>>> )
torch.utils.checkpoint.create_selective_checkpoint_contexts(policy_fn_or_list, allow_cache_entry_mutation=False) [исходный код]

Вспомогательная функция, предотвращающая повторное вычисление некоторых операций при создании контрольных точек активаций.

Используйте эту функцию вместе с torch.utils.checkpoint.checkpoint, чтобы управлять тем, какие операции повторно вычисляются во время обратного прохода.

Параметры:
  • policy_fn_or_list (Callable или List) –

    • Если передана функция политики, она должна принимать SelectiveCheckpointContext, OpOverload, аргументы args и kwargs операции и возвращать значение перечисления CheckpointPolicy, указывающее, следует ли повторно вычислять операцию.
    • Если передан список операций, это эквивалентно политике, возвращающей CheckpointPolicy.MUST_SAVE для указанных операций и CheckpointPolicy.PREFER_RECOMPUTE для всех остальных.
  • allow_cache_entry_mutation (bool, необязательный) – По умолчанию возникает ошибка, если изменяются какие-либо тензоры, кэшированные выборочной контрольной точкой активаций, чтобы обеспечить корректность. Если задано значение True, эта проверка отключается.
Возвращает:

Кортеж из двух менеджеров контекста.

Пример

>>> import functools
>>>
>>> x = torch.rand(10, 10, requires_grad=True)
>>> y = torch.rand(10, 10, requires_grad=True)
>>>
>>> ops_to_save = [
>>>    torch.ops.aten.mm.default,
>>> ]
>>>
>>> def policy_fn(ctx, op, *args, **kwargs):
>>>    if op in ops_to_save:
>>>        return CheckpointPolicy.MUST_SAVE
>>>    else:
>>>        return CheckpointPolicy.PREFER_RECOMPUTE
>>>
>>> context_fn = functools.partial(create_selective_checkpoint_contexts, policy_fn)
>>>
>>> # or equivalently
>>> context_fn = functools.partial(create_selective_checkpoint_contexts, ops_to_save)
>>>
>>> def fn(x, y):
>>>     return torch.sigmoid(torch.matmul(torch.matmul(x, y), y)) * y
>>>
>>> out = torch.utils.checkpoint.checkpoint(
>>>     fn, x, y,
>>>     use_reentrant=False,
>>>     context_fn=context_fn,
>>> )
class torch.utils.checkpoint.GraphExecGroup [исходный код]

Все области с контрольными точками, встречающиеся при обратном проходе внутри одного экземпляра этого менеджера контекста, будут повторно вычислены не более одного раза, даже если обратный проход вызывается несколько раз.

Вызовы обратного прохода внутри одного экземпляра этого менеджера контекста должны выполняться по непересекающимся областям графа обратного прохода, даже если retain_graph=True. В частности, два вызова обратного прохода не могут использовать одну и ту же сохранённую активацию для вычисления градиентов.

Примечание

Этот менеджер контекста влияет только на контрольные точки с use_reentrant=False и в остальных случаях ничего не делает.

torch.utils.checkpoint.set_checkpoint_early_stop(enable) [исходный код]

Управляет тем, следует ли контрольной точке досрочно прекращать повторное вычисление.

По умолчанию нереентерабельная контрольная точка прекращает повторное вычисление, как только вычислит все необходимые тензоры. Если эта возможность мешает работе вашего приложения, её можно отключить с помощью этого менеджера контекста.

Этот менеджер контекста должен быть активен только во время прямого прохода. Во время обратного прохода он не требуется.

Пример:

>>> message = "saved tensors default hooks are disabled"
>>> with set_checkpoint_early_stop(False):
...     # Any checkpoint under this context manager will respect this
...     # context manager, even if its backward is performed outside.
...     out = checkpoint(fn, inputs)
...
>>> out.backward()
torch.utils.checkpoint.set_device_states(devices, states, *, device_type=None) [исходный код]

Задаёт состояния генератора случайных чисел для указанных устройств.

Параметры:
  • devices – идентификаторы устройств, для которых задаются состояния.
  • states – задаваемые состояния.
  • device_type – device_type устройств, для которых задаются состояния. По умолчанию используется устройство, возвращаемое вызовом DefaultDeviceType.get_device_type(); это cuda, если значение не изменено вызовом DefaultDeviceType::set_device_type().

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/checkpoint.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API