Spec-Zone.ru › PyTorch 1

torch.utils.checkpoint

Примечание

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

Логика сохранения сохраняет и восстанавливает состояние RNG для текущего устройства и устройства всех аргументов тензора CUDA в run_fn. Однако эта логика не может предвидеть, переместит ли пользователь тензоры на новое устройство в рамках run_fn. Поэтому, если вы перемещаете тензоры на новое устройство («новое» означает не принадлежащее набору [текущее устройство + устройства аргументов тензора]) внутри run_fn, детерминированный вывод по сравнению с некэшированными проходами никогда не гарантируется.

torch.utils.checkpoint.checkpoint(function, *args, use_reentrant=True, **kwargs) [source]

Кэширование модели или части модели

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

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

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

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

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

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

Если use_reentrant=True указан, если кэшированный сегмент содержит тензоры, открепленные от вычислительной графы с помощью detach() или torch.no_grad(), обратный проход выдаст ошибку. Это происходит потому, что checkpoint заставляет все выходы требовать градиенты, что вызывает проблемы, когда для тензора в модели определено отсутствие градиента. Чтобы обойти это, открепите тензоры за пределами функции checkpoint. Обратите внимание, что кэшированный сегмент может содержать тензоры, открепленные от вычислительной графы, если use_reentrant=False указан.

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

Если use_reentrant=True указан, по крайней мере один из входных данных должен иметь requires_grad=True, если нужны градиенты для входных данных модели, в противном случае кэшированная часть модели не будет иметь градиенты. По крайней мере один из выходов также должен иметь requires_grad=True. Обратите внимание, что это не относится, если use_reentrant=False указан.

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

Если use_reentrant=True указан, кэширование в настоящее время поддерживает только torch.autograd.backward(), и только если аргумент inputs не передан. torch.autograd.grad() не поддерживается. Если use_reentrant=False указан, кэширование будет работать с torch.autograd.grad().

Параметры:
  • function – описывает то, что нужно выполнить в прямом проходе модели или части модели. Он также должен знать, как обрабатывать переданные как кортеж входные данные. Например, в LSTM, если пользователь передает (activation, hidden), function должен правильно использовать первый вход как activation и второй вход как hidden
  • preserve_rng_state (bool, необязательно) – Пропустить сохранение и восстановление состояния RNG во время каждого кэша. Значение по умолчанию: True
  • use_reentrant (bool, необязательно) – Использовать реализацию кэширования, требующую рекурсивного autograd. Если use_reentrant=False указан, checkpoint будет использовать реализацию, не требующую рекурсивного autograd. Это позволяет checkpoint поддерживать дополнительную функциональность, например, работать как ожидается с torch.autograd.grad и поддержкой ключевых аргументов входных данных в кэшированную функцию. Обратите внимание, что в будущих версиях PyTorch по умолчанию будет use_reentrant=False. Значение по умолчанию: True
  • args – кортеж, содержащий входные данные для function
Возвращает:

Вывод выполнения function на *args

torch.utils.checkpoint.checkpoint_sequential(functions, segments, input, **kwargs) [source]

Вспомогательная функция для кэширования последовательных моделей.

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

См. checkpoint() о том, как работает кэширование.

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

Кэширование в настоящее время поддерживает только torch.autograd.backward(), и только если аргумент inputs не передан. torch.autograd.grad() не поддерживается.

Параметры:
  • functions – torch.nn.Sequential или список модулей или функций (составляющих модель), которые нужно выполнить последовательно.
  • segments – Количество фрагментов для создания в модели
  • input – Тензор, являющийся входом для functions
  • preserve_rng_state (bool, необязательно) – Пропустить сохранение и восстановление состояния RNG во время каждого кэша. Значение по умолчанию: True
Возвращает:

Вывод выполнения functions последовательно на *inputs

Пример

>>> model = nn.Sequential(...)
>>> input_var = checkpoint_sequential(model, chunks, input_var)

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

Spec-Zone.ru

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