Адам
-
class torch.optim.Adam(params, lr=0.001, betas=(0.9, 0.999), eps=1e-08, weight_decay=0, amsgrad=False, *, foreach=None, maximize=False, capturable=False, differentiable=False, fused=None)[source] -
Реализует алгоритм Adam.
Для получения дополнительной информации об алгоритме мы обращаемся к Adam: A Method for Stochastic Optimization.
- Параметры
-
- params (iterable) – итерируемый список параметров для оптимизации или словари, определяющие группы параметров
- lr (float, Tensor, необязательно) – скорость обучения (по умолчанию: 1e-3). Скорость обучения в виде тензора пока не поддерживается во всех наших реализациях. Используйте float LR, если вы не задаёте также fused=True или capturable=True.
- betas (Tuple[float, float], необязательно) – коэффициенты, используемые для вычисления скользящих средних градиента и его квадрата (по умолчанию: (0.9, 0.999))
- eps (float, необязательно) – слагаемое, добавляемое в знаменатель для улучшения числовой устойчивости (по умолчанию: 1e-8)
- weight_decay (float, необязательно) – коэффициент затухания (штраф L2) (по умолчанию: 0)
- amsgrad (bool, необязательно) – использовать ли вариант AMSGrad этого алгоритма из статьи On the Convergence of Adam and Beyond (по умолчанию: False)
- foreach (bool, необязательно) – использовать ли реализацию оптимизатора foreach. Если пользователь не указывает (то есть foreach равно None), мы будем пытаться использовать foreach вместо цикла for на CUDA, поскольку обычно это значительно производительнее. Обратите внимание, что реализация foreach использует ~ sizeof(params) больше пиковой памяти, чем версия цикла for, из-за того, что промежуточные результаты являются списком тензоров, а не просто одним тензором. Если память ограничена, передавайте меньше параметров в оптимизатор за раз или установите этот флаг в значение False (по умолчанию: None)
- maximize (bool, необязательно) – максимизировать параметры на основе целевой функции вместо минимизации (по умолчанию: False)
- capturable (bool, необязательно) – является ли данный экземпляр подходящим для захвата в графе CUDA. Передача True может ухудшить производительность без захвата, поэтому если вы не планируете захватывать этот экземпляр, оставьте False (по умолчанию: False)
- differentiable (bool, необязательно) – нужно ли использовать autograd при шаге оптимизатора в процессе обучения. В противном случае функция step() выполняется в контексте torch.no_grad(). Установка в значение True может ухудшить производительность, поэтому оставьте False, если вы не планируете использовать autograd с этим экземпляром (по умолчанию: False)
-
fused (bool, необязательно) – использовать ли объединённую реализацию (только CUDA). В настоящее время поддерживаются
torch.float64,torch.float32,torch.float16, иtorch.bfloat16(по умолчанию: None)
Примечание
Реализации foreach и fused обычно быстрее, чем реализация с циклом for и одним тензором. Поэтому, если пользователь не указал ОБА флага (т. е. когда foreach = fused = None), мы попытаемся по умолчанию использовать реализацию foreach, когда все тензоры находятся на CUDA. Например, если пользователь указывает True для fused, но ничего для foreach, мы будем выполнять объединённую реализацию. Если пользователь указывает False для foreach, но ничего для fused (или False для fused, но ничего для foreach), мы будем выполнять реализацию с циклом for. Если пользователь указывает True для foreach и fused, мы будем отдавать предпочтение fused, так как она обычно быстрее. Мы пытаемся использовать самую быструю, поэтому иерархия идёт fused -> foreach -> цикл for. ОДНАКО, поскольку объединённая реализация относительно новая, мы хотим дать ей достаточное время на отработку, поэтому по умолчанию мы используем 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. Хук может изменить state_dict на месте или, необязательно, вернуть новый. Если возвращается 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) -
Регистрация обработчика состояния словаря после вызова
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) -
Регистрация обработчика состояния словаря перед вызовом
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] } ] } -
-
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), за которым следует обратный проход, градиенты для параметров, которые не получили градиент, гарантированно будут 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.Adam.html