AdamW
-
class torch.optim.AdamW(params, lr=0.001, betas=(0.9, 0.999), eps=1e-08, weight_decay=0.01, amsgrad=False, *, maximize=False, foreach=None, capturable=False, differentiable=False, fused=None)[source] -
Реализует алгоритм AdamW.
Для получения дополнительной информации об алгоритме, мы обращаем вас к Decoupled Weight Decay Regularization.
- Параметры
-
- params (iterable) – iterable параметров для оптимизации или словари, определяющие группы параметров
- lr (float, Тензор, необязательно) – скорость обучения (по умолчанию: 1e-3). Тензор LR пока не поддерживается во всех наших реализациях. Используйте float LR, если вы не задаёте также fused=True или capturable=True.
- betas (Tuple[float, float], необязательно) – коэффициенты, используемые для вычисления усреднённых градиентов и его квадрата (по умолчанию: (0.9, 0.999))
- eps (float, необязательно) – член, добавляемый в знаменатель для повышения числовой устойчивости (по умолчанию: 1e-8)
- weight_decay (float, необязательно) – коэффициент распада веса (по умолчанию: 1e-2)
- amsgrad (bool, необязательно) – использовать ли вариант AMSGrad этого алгоритма из статьи On the Convergence of Adam and Beyond (по умолчанию: False)
- maximize (bool, необязательно) – максимизировать параметры на основе целевой функции вместо минимизации (по умолчанию: False)
- foreach (bool, необязательно) – использовать ли foreach-реализацию оптимизатора. Если пользователь не указал (т.е., foreach = None), мы будем пытаться использовать foreach вместо реализации циклом по CUDA, так как она обычно значительно производительнее. Обратите внимание, что foreach-реализация использует ~ sizeof(params) больше пиковой памяти, чем версия с циклом, из-за того, что промежуточные данные являются tensorlist, а не просто одним тензором. Если память ограничена, передавайте меньше параметров в оптимизатор за раз или установите этот флаг в False (по умолчанию: None)
- capturable (bool, необязательно) – можно ли захватить этот экземпляр в график CUDA. Передача True может ухудшить производительность вне графика, поэтому, если вы не планируете захват этого экземпляра, оставьте False (по умолчанию: False)
- differentiable (bool, необязательно) – нужно ли проводить автодифференцирование через шаг оптимизатора при обучении. В противном случае функция step() выполняется в контексте torch.no_grad(). Установка в True может ухудшить производительность, поэтому оставьте False, если вы не планируете проводить автодифференцирование через этот экземпляр (по умолчанию: False)
-
fused (bool, необязательно) – использовать ли объединённую реализацию (только CUDA). В настоящее время поддерживаются
torch.float64,torch.float32,torch.float16, иtorch.bfloat16. (по умолчанию: None)
Примечание
Реализации foreach и fused обычно быстрее, чем реализация циклом по одному тензору. Поэтому, если пользователь не указал ОБА флага (т.е., когда foreach = fused = None), мы попытаемся по умолчанию использовать foreach-реализацию, когда все тензоры находятся на CUDA. Например, если пользователь указывает True для fused, но ничего для foreach, мы будем выполнять fused-реализацию. Если пользователь указывает False для foreach, но ничего для fused (или False для fused, но ничего для foreach), мы будем выполнять реализацию циклом. Если пользователь указывает True для обоих foreach и fused, мы будем отдавать предпочтение fused, так как она обычно быстрее. Мы пытаемся использовать самую быструю реализацию, поэтому иерархия выглядит так: fused -> foreach -> цикл. Однако, поскольку реализация fused относительно новая, мы хотим дать ей достаточное время для проверки, поэтому мы по умолчанию используем foreach, а НЕ fused, когда пользователь не указал ни один из этих флагов.
-
add_param_group(param_group) -
Добавить группу параметров в
Optimizerparam_groups.Это может быть полезно при донастройке предварительно обученной сети, так как замороженные слои можно сделать обучаемыми и добавить к
Optimizerпо мере обучения.- Параметры
-
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, предоставленный пост-хук будет вызван до всех уже зарегистрированных пост-хуков на
load_state_dict. В противном случае, предоставленный пост-хук будет вызван после всех уже зарегистрированных пост-хуков. (по умолчанию: 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. Хук может изменять состояние словаря на месте или, по желанию, возвращать новый. Если возвращается словарь состояния, он будет загружен в оптимизатор.Хук будет вызван с аргументами
selfиstate_dictперед вызовомload_state_dictнаself. Зарегистрированный хук может использоваться для предварительной обработки перед тем, как будет сделан вызовload_state_dict.- Параметры
-
- hook (Callable) – Пользовательский хук для регистрации.
-
prepend (bool) – Если True, предоставленный пре-хук будет вызван до всех уже зарегистрированных пре-хуков на
load_state_dict. В противном случае, предоставленный пре-хук будет вызван после всех уже зарегистрированных пре-хуков. (по умолчанию: False)
- Возвращает
-
дескриптор, который можно использовать для удаления добавленного хука, вызвав
handle.remove() - Тип возвращаемого значения
-
torch.utils.hooks.RemoveableHandle
-
register_state_dict_post_hook(hook, prepend=False) -
Регистрация обработчика post-состояния словаря, который будет вызван после вызова
state_dict(). Он должен иметь следующий сигнатуру:hook(optimizer, state_dict) -> state_dict or None
Обработчик будет вызван с аргументами
selfиstate_dictпосле генерацииstate_dictдляself. Обработчик может изменять словарь состояния на месте или, по желанию, возвращать новый. Зарегистрированный обработчик можно использовать для выполнения пост-обработки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) -
Регистрация обработчика pre-состояния словаря, который будет вызван до вызова
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 модифицируются пре-обработчиком, то преобразованные значения возвращаются в виде кортежа, содержащего новые значения args и 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 -
параметр группы — это словарь. Каждая группа параметров содержит метаданные, специфичные для оптимизатора, такие как скорость обучения и коэффициент затухания, а также список идентификаторов параметров параметров в группе.
-
ПРИМЕЧАНИЕ: Идентификаторы параметров могут выглядеть как индексы, но они представляют собой просто идентификаторы, связывающие состояние с группой параметров. При загрузке из словаря состояния оптимизатор свяжет группу параметров
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] } ] } -
-
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.AdamW.html