RMSprop
-
class torch.optim.RMSprop(params, lr=0.01, alpha=0.99, eps=1e-08, weight_decay=0, momentum=0, centered=False, foreach=None, maximize=False, differentiable=False)[source] -
Реализует алгоритм RMSprop.
Для получения дополнительной информации об алгоритме см. лекции Дж. Хинтона. и централизованную версию Генерация последовательностей с рекуррентными нейронными сетями. В данном случае реализация извлекает квадратный корень из среднего градиента перед добавлением ε (обратите внимание, что TensorFlow меняет эти две операции). Таким образом, эффективная скорость обучения составляет , где — запланированная скорость обучения, а — взвешенное скользящее среднее квадрата градиента.
- Параметры
-
- params (iterable) – перечисляемый объект параметров для оптимизации или словари, определяющие группы параметров
- lr (float, необязательно) – скорость обучения (по умолчанию: 1e-2)
- momentum (float, необязательно) – коэффициент импульса (по умолчанию: 0)
- alpha (float, необязательно) – константа сглаживания (по умолчанию: 0.99)
- eps (float, необязательно) – член, добавляемый к знаменателю для улучшения числовой устойчивости (по умолчанию: 1e-8)
-
centered (bool, необязательно) – если
True, вычисляется центрированный RMSProp, градиент нормализуется по оценке его дисперсии - weight_decay (float, необязательно) – сжатие весов (штраф L2) (по умолчанию: 0)
- foreach (bool, необязательно) – использовать ли реализацию оптимизатора foreach. Если пользователь не указывает (т.е. foreach равно None), мы попытаемся использовать foreach вместо реализации циклом for на CUDA, так как это обычно значительно производительнее. Обратите внимание, что реализация foreach использует ~ sizeof(params) больше пиковой памяти, чем версия с циклом for, из-за того, что промежуточные результаты являются списком тензоров, а не просто одним тензором. Если память ограничена, передавайте меньше параметров через оптимизатор за раз или установите этот флаг в False (по умолчанию: None)
- maximize (bool, необязательно) – максимизировать параметры на основе целевой функции вместо минимизации (по умолчанию: False)
- differentiable (bool, необязательно) – должен ли автоград происходить через шаг оптимизатора при обучении. В противном случае функция step() выполняется в контексте torch.no_grad(). Установка в True может ухудшить производительность, поэтому оставьте False, если вы не планируете выполнять автоград через этот экземпляр (по умолчанию: False)
-
add_param_group(param_group) -
Добавить группу параметров к
Optimizersparam_groups.Это может быть полезно при донастройке предварительно обученной сети, поскольку замороженные слои можно сделать обучаемыми и добавить к
Optimizerпо мере обучения.- Параметры
-
param_group (dict) – Указывает, какие тензоры должны быть оптимизированы вместе со специфичными для группы параметрами оптимизации.
-
load_state_dict(state_dict) -
Загружает состояние оптимизатора.
- Параметры
-
state_dict (dict) – состояние оптимизатора. Должно быть объектом, возвращённым из вызова
state_dict().
-
register_load_state_dict_post_hook(hook, prepend=False) -
Регистрирует обработчик post-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) -
Регистрирует обработчик pre-load_state_dict, который вызывается до вызова
load_state_dict(). Он должен иметь следующий вид:hook(optimizer, state_dict) -> state_dict or None
Аргумент
optimizer— это экземпляр оптимизатора, используемый, а аргументstate_dict— это поверхностная копияstate_dict, переданная пользователем вload_state_dict. Обработчик может изменять state_dict inplace или, по желанию, возвращать новый. Если возвращается 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, который будет вызван после вызова
state_dict(). Он должен иметь следующий вид:hook(optimizer, state_dict) -> state_dict or None
Обработчик вызывается с аргументами
selfиstate_dictпосле генерацииstate_dictнаself. Обработчик может изменять state_dict inplace или, по желанию, возвращать новый. Зарегистрированный обработчик может использоваться для выполнения пост-обработки наstate_dictперед возвратом.- Параметры
-
- hook (Callable) – пользовательский обработчик, который необходимо зарегистрировать.
-
prepend (bool) – Если True, предоставленный пост-обработчик будет вызван до всех уже зарегистрированных пост-обработчиков на
state_dict. В противном случае предоставленный пост-обработчик будет вызван после всех уже зарегистрированных пост-обработчиков. (по умолчанию: 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.RMSprop.html