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 – описывает то, что нужно выполнить в прямом проходе модели или части модели. Он также должен знать, как обрабатывать переданные как кортеж входные данные. Например, в LSTM, если пользователь передает
- Возвращает:
-
Вывод выполнения
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 –
- Возвращает:
-
Вывод выполнения
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