Spec-Zone.ru › PyTorch 2

NAdam

class torch.optim.NAdam(params, lr=0.002, betas=(0.9, 0.999), eps=1e-08, weight_decay=0, momentum_decay=0.004, decoupled_weight_decay=False, *, foreach=None, capturable=False, differentiable=False) [source]

Реализует алгоритм NAdam.

input:γt (lr),β1,β2 (betas),θ0 (params),f(θ) (objective)λ (weight decay),ψ (momentum decay)decoupled_weight_decayinitialize:m0←0 ( first moment),v0←0 ( second moment)fort=1to…dogt←∇θft(θt−1)θt←θt−1ifλ≠0ifdecoupled_weight_decayθt←θt−1−γλθt−1elsegt←gt+λθt−1μt←β1(1−120.96tψ)μt+1←β1(1−120.96(t+1)ψ)mt←β1mt−1+(1−β1)gtvt←β2vt−1+(1−β2)gt2mt^←μt+1mt/(1−∏i=1t+1μi)+(1−μt)gt/(1−∏i=1tμi)vt^←vt/(1−β2t)θt←θt−γmt^/(vt^+ϵ)returnθt\begin{aligned} &\rule{110mm}{0.4pt} \\ &\textbf{input} : \gamma_t \text{ (lr)}, \: \beta_1,\beta_2 \text{ (betas)}, \: \theta_0 \text{ (params)}, \: f(\theta) \text{ (objective)} \\ &\hspace{13mm} \: \lambda \text{ (weight decay)}, \:\psi \text{ (momentum decay)} \\ &\hspace{13mm} \: \textit{decoupled\_weight\_decay} \\ &\textbf{initialize} : m_0 \leftarrow 0 \text{ ( first moment)}, v_0 \leftarrow 0 \text{ ( second moment)} \\[-1.ex] &\rule{110mm}{0.4pt} \\ &\textbf{for} \: t=1 \: \textbf{to} \: \ldots \: \textbf{do} \\ &\hspace{5mm}g_t \leftarrow \nabla_{\theta} f_t (\theta_{t-1}) \\ &\hspace{5mm} \theta_t \leftarrow \theta_{t-1} \\ &\hspace{5mm} \textbf{if} \: \lambda \neq 0 \\ &\hspace{10mm}\textbf{if} \: \textit{decoupled\_weight\_decay} \\ &\hspace{15mm} \theta_t \leftarrow \theta_{t-1} - \gamma \lambda \theta_{t-1} \\ &\hspace{10mm}\textbf{else} \\ &\hspace{15mm} g_t \leftarrow g_t + \lambda \theta_{t-1} \\ &\hspace{5mm} \mu_t \leftarrow \beta_1 \big(1 - \frac{1}{2} 0.96^{t \psi} \big) \\ &\hspace{5mm} \mu_{t+1} \leftarrow \beta_1 \big(1 - \frac{1}{2} 0.96^{(t+1)\psi}\big)\\ &\hspace{5mm}m_t \leftarrow \beta_1 m_{t-1} + (1 - \beta_1) g_t \\ &\hspace{5mm}v_t \leftarrow \beta_2 v_{t-1} + (1-\beta_2) g^2_t \\ &\hspace{5mm}\widehat{m_t} \leftarrow \mu_{t+1} m_t/(1-\prod_{i=1}^{t+1}\mu_i)\\[-1.ex] & \hspace{11mm} + (1-\mu_t) g_t /(1-\prod_{i=1}^{t} \mu_{i}) \\ &\hspace{5mm}\widehat{v_t} \leftarrow v_t/\big(1-\beta_2^t \big) \\ &\hspace{5mm}\theta_t \leftarrow \theta_t - \gamma \widehat{m_t}/ \big(\sqrt{\widehat{v_t}} + \epsilon \big) \\ &\rule{110mm}{0.4pt} \\[-1.ex] &\bf{return} \: \theta_t \\[-1.ex] &\rule{110mm}{0.4pt} \\[-1.ex] \end{aligned}

Для получения дополнительной информации об алгоритме мы обращаем вас к Incorporating Nesterov Momentum into Adam.

Параметры
  • params (iterable) – iterable параметров для оптимизации или словари, определяющие группы параметров
  • lr (float, необязательно) – скорость обучения (по умолчанию: 2e-3)
  • betas (Tuple[float, float], необязательно) – коэффициенты, используемые для вычисления скользящих средних градиента и его квадрата (по умолчанию: (0.9, 0.999))
  • eps (float, необязательно) – член, добавляемый к знаменателю для повышения числовой устойчивости (по умолчанию: 1e-8)
  • weight_decay (float, необязательно) – регуляризация по весам (штраф L2) (по умолчанию: 0)
  • momentum_decay (float, необязательно) – коэффициент затухания импульса (по умолчанию: 4e-3)
  • decoupled_weight_decay (bool, необязательно) – использовать ли разъединенную регуляризацию весов, как в AdamW, чтобы получить NAdamW (по умолчанию: False)
  • foreach (bool, необязательно) – использовать ли foreach реализацию оптимизатора. Если пользователь не указывает, мы будем пытаться использовать foreach на CUDA, так как это обычно значительно производительнее. Обратите внимание, что foreach реализация использует ~ sizeof(params) больше пиковой памяти, чем версия for-loop, из-за того, что промежуточные значения являются тензорным списком, а не одним тензором. Если память ограничена, обрабатывайте меньше параметров через оптимизатор за раз или установите этот флаг в False (по умолчанию: None)
  • capturable (bool, необязательно) – можно ли захватить этот экземпляр в графе CUDA. Передача True может ухудшить производительность без графа, поэтому, если вы не собираетесь захватывать этот экземпляр, оставьте False (по умолчанию: False)
  • differentiable (bool, необязательно) – должен ли autograd происходить через шаг оптимизатора в процессе обучения. В противном случае, функция step() выполняется в контексте torch.no_grad(). Установка в True может ухудшить производительность, поэтому оставьте False, если вы не хотите запускать autograd через этот экземпляр (по умолчанию: False)
add_param_group(param_group)

Добавить группу параметров в оптимизаторы param_groups.

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

Параметры

param_group (dict) – Указывает, какие тензоры должны быть оптимизированы вместе с групповыми параметрами оптимизации.

load_state_dict(state_dict)

Загружает состояние оптимизатора.

Параметры

state_dict (dict) – состояние оптимизатора. Должно быть объектом, возвращённым из вызова state_dict().

register_load_state_dict_post_hook(hook, prepend=False)

Регистрирует пост-обработчик load_state_dict, который будет вызван после вызова load_state_dict(). Он должен иметь следующий вид:

hook(optimizer) -> None

Аргумент optimizer — это экземпляр оптимизатора, используемый.

Обработчик вызывается с аргументом self после вызова load_state_dict на self. Зарегистрированный обработчик может использоваться для выполнения пост-обработки после того, как load_state_dict загрузил state_dict.

Параметры
  • hook (Callable) – Зарегистрированный пользовательский обработчик.
  • prepend (bool) – Если True, предоставленный пост hook будет запущен до всех уже зарегистрированных пост-обработчиков на load_state_dict. В противном случае, предоставленный hook будет запущен после всех уже зарегистрированных пост-обработчиков. (по умолчанию: False)
Возвращает

дескриптор, который может быть использован для удаления добавленного обработчика, вызвав handle.remove()

Тип возвращаемого значения

torch.utils.hooks.RemoveableHandle

register_load_state_dict_pre_hook(hook, prepend=False)

Регистрирует пре-обработчик load_state_dict, который будет вызван до вызова load_state_dict(). Он должен иметь следующий вид:

hook(optimizer, state_dict) -> state_dict or None

Аргумент optimizer — это экземпляр оптимизатора, используемый, а аргумент state_dict — это поверхностная копия state_dict , переданная пользователем в load_state_dict . Обработчик может изменить state_dict на месте или, по желанию, вернуть новый. Если возвращается state_dict, он будет загружен в оптимизатор.

Обработчик вызывается с аргументом self и state_dict перед вызовом load_state_dict на self. Зарегистрированный обработчик может использоваться для выполнения предварительной обработки перед вызовом load_state_dict.

Параметры
  • hook (Callable) – Зарегистрированный пользовательский обработчик.
  • prepend (bool) – Если True, предоставленный пре hook будет запущен до всех уже зарегистрированных пре-обработчиков на load_state_dict. В противном случае, предоставленный hook будет запущен после всех уже зарегистрированных пре-обработчиков. (по умолчанию: False)
Возвращает

дескриптор, который может быть использован для удаления добавленного обработчика, вызвав handle.remove()

Тип возвращаемого значения

torch.utils.hooks.RemoveableHandle

register_state_dict_post_hook(hook, prepend=False)

Регистрирует пост-обработчик state_dict, который будет вызван после вызова state_dict(). Он должен иметь следующий вид:

hook(optimizer, state_dict) -> state_dict or None

Обработчик вызывается с аргументами self и state_dict после создания state_dict на self. Обработчик может изменить state_dict на месте или, по желанию, вернуть новый. Зарегистрированный обработчик может использоваться для выполнения пост-обработки над state_dict перед возвратом.

Параметры
  • hook (Callable) – Зарегистрированный пользовательский обработчик.
  • prepend (bool) – Если True, предоставленный пост hook будет запущен до всех уже зарегистрированных пост-обработчиков на state_dict. В противном случае, предоставленный hook будет запущен после всех уже зарегистрированных пост-обработчиков. (по умолчанию: False)
Возвращает

дескриптор, который может быть использован для удаления добавленного обработчика, вызвав handle.remove()

Тип возвращаемого значения

torch.utils.hooks.RemoveableHandle

register_state_dict_pre_hook(hook, prepend=False)

Регистрация предварительного хука state dict, который будет вызван перед вызовом state_dict(). Он должен иметь следующий вид:

hook(optimizer) -> None

Аргумент optimizer — это экземпляр оптимизатора, используемый. Хук будет вызван с аргументом self перед вызовом state_dict для self. Зарегистрированный хук может быть использован для выполнения предобработки перед вызовом state_dict.

Параметры
  • hook (Callable) – Пользовательский хук, который нужно зарегистрировать.
  • prepend (bool) – Если True, предоставленный предварительный hook будет вызван перед всеми уже зарегистрированными предварительными хуками для state_dict. В противном случае, предоставленный hook будет вызван после всех уже зарегистрированных предварительных хуков. (по умолчанию: False)
Возвращаемое значение

дескриптор, который можно использовать для удаления добавленного хука, вызвав handle.remove()

Тип возвращаемого значения

torch.utils.hooks.RemoveableHandle

register_step_post_hook(hook)

Регистрация хука после шага оптимизатора, который будет вызван после шага оптимизации. Он должен иметь следующий вид:

hook(optimizer, args, kwargs) -> None

Аргумент optimizer — это экземпляр оптимизатора, используемый.

Параметры

hook (Callable) – Пользовательский хук, который нужно зарегистрировать.

Возвращаемое значение

дескриптор, который можно использовать для удаления добавленного хука, вызвав handle.remove()

Тип возвращаемого значения

torch.utils.hooks.RemovableHandle

register_step_pre_hook(hook)

Регистрация хука перед шагом оптимизатора, который будет вызван перед шагом оптимизации. Он должен иметь следующий вид:

hook(optimizer, args, kwargs) -> None or modified args and kwargs

Аргумент optimizer — это экземпляр оптимизатора, используемый. Если значения args и kwargs изменяются предварительным хуком, то преобразованные значения возвращаются как кортеж, содержащий новые значения new_args и new_kwargs.

Параметры

hook (Callable) – Пользовательский хук, который нужно зарегистрировать.

Возвращаемое значение

дескриптор, который можно использовать для удаления добавленного хука, вызвав handle.remove()

Тип возвращаемого значения

torch.utils.hooks.RemovableHandle

state_dict()

Возвращает состояние оптимизатора в виде dict.

Он содержит две записи:

  • state: a Dict holding current optimization state. Its content

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

  • param_groups: a List containing all parameter groups where each

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

ПРИМЕЧАНИЕ: Идентификаторы параметров могут выглядеть как индексы, но они просто идентификаторы, связывающие состояние с param_group. При загрузке из state_dict оптимизатор объединит param_group params (целочисленные идентификаторы) и оптимизатор param_groups (фактические nn.Parameter ) для соответствия состоянию БЕЗ дополнительной проверки.

Возвращаемый словарь состояния может выглядеть примерно так:

{
    'state': {
        0: {'momentum_buffer': tensor(...), ...},
        1: {'momentum_buffer': tensor(...), ...},
        2: {'momentum_buffer': tensor(...), ...},
        3: {'momentum_buffer': tensor(...), ...}
    },
    'param_groups': [
        {
            'lr': 0.01,
            'weight_decay': 0,
            ...
            'params': [0]
        },
        {
            'lr': 0.001,
            'weight_decay': 0.5,
            ...
            'params': [1, 2, 3]
        }
    ]
}
Тип возвращаемого значения

Dict[str, Any]

step(closure=None) [source]

Выполняет один шаг оптимизации.

Параметры

closure (Callable, optional) – Закрытие, которое повторно оценивает модель и возвращает потерю.

zero_grad(set_to_none=True)

Сбрасывает градиенты всех оптимизированных torch.Tensor.

Параметры

set_to_none (bool) – вместо присвоения нуля, установить градиенты в None. Это, как правило, приведет к меньшему объёму памяти и может незначительно улучшить производительность. Однако это меняет определённое поведение. Например: 1. Когда пользователь пытается получить доступ к градиенту и выполнить на нём ручные операции, атрибут None или тензор, заполненный нулями, будет вести себя по-разному. 2. Если пользователь запрашивает zero_grad(set_to_none=True), за которым следует обратный проход, градиенты .grad гарантированно будут None для параметров, которые не получили градиент. 3. torch.optim оптимизаторы ведут себя по-разному, если градиент равен 0 или None (в одном случае он выполняет шаг с градиентом 0, а в другом случае пропускает шаг).

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.optim.NAdam.html

Spec-Zone.ru

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