Spec-Zone.ru › PyTorch 2

torch.utils.checkpoint

Примечание

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

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

torch.utils.checkpoint.checkpoint(function, *args, use_reentrant=None, context_fn=<function noop_context_fn>, determinism_check='default', debug=False, **kwargs) [source]

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

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

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

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

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

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

Если вы используете вариант 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
  • preserve_rng_state (bool, optional) – пропускает сохранение и восстановление состояния генератора случайных чисел во время каждого кеширования. Значение по умолчанию: True
  • use_reentrant (bool, optional) – использовать реализацию кеширования, которая требует рекурсивного autograd. Если указано use_reentrant=False, checkpoint будет использовать реализацию, которая не требует рекурсивного autograd. Это позволяет checkpoint поддерживать дополнительные функции, например, работать как ожидается с torch.autograd.grad и поддерживать входные ключевые аргументы в кешированную функцию. Обратите внимание, что в будущих версиях PyTorch значение по умолчанию будет use_reentrant=False. Значение по умолчанию: True
  • context_fn (Callable, optional) – функция, возвращающая кортеж из двух менеджеров контекста. Функция и ее перевычисление будут выполняться под управлением первого и второго менеджеров контекста соответственно. Этот аргумент поддерживается только если use_reentrant=False.
  • determinism_check (str, optional) – строка, определяющая проверку детерминизма. По умолчанию она устанавливается в "default", которая сравнивает формы, типы данных и устройства перевычисленных тензоров с сохраненными тензорами. Чтобы отключить эту проверку, укажите "none". В настоящее время это единственные два поддерживаемых значения. Откройте вопрос, если вы хотите увидеть больше проверок детерминизма. Этот аргумент поддерживается только если use_reentrant=False, если use_reentrant=True, проверка детерминизма всегда отключена.
  • debug (bool, optional) – Если True, сообщения об ошибках также будут содержать трассировку операторов, запущенных во время исходного вычисления прямого прохода, а также перевычисления. Этот аргумент поддерживается только если use_reentrant=False.
  • args – кортеж, содержащий входные данные для function
Возвращает

Результат выполнения function на *args

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

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

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

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

Если вы используете use_reentrant=True` variant (this is the default), 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, optional) – Пропустить сохранение и восстановление состояния генератора случайных чисел во время каждого кэширования. По умолчанию: True
  • use_reentrant (bool, optional) – Использовать реализацию кэширования, которая требует рекурсивного autograd. Если use_reentrant=False указано, checkpoint будет использовать реализацию, которая не требует рекурсивного autograd. Это позволяет checkpoint поддерживать дополнительную функциональность, например, работать как ожидается с torch.autograd.grad и поддерживать ключевые аргументы, передаваемые в кэшированную функцию. По умолчанию: 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/2.1/checkpoint.html

Spec-Zone.ru

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