SGD
-
class torch.optim.SGD(params, lr=<required parameter>, momentum=0, dampening=0, weight_decay=0, nesterov=False, *, maximize=False, foreach=None, differentiable=False)[source] -
Реализует стохастический градиентный спуск (по желанию с моментом).
Момент Нестерова основан на формуле из На важности инициализации и момента в глубоком обучении.
- Параметры
-
- params (iterable) – iterable параметров для оптимизации или словари, определяющие группы параметров
- lr (float) – скорость обучения
- momentum (float, optional) – коэффициент момента (по умолчанию: 0)
- weight_decay (float, optional) – спад веса (штраф L2) (по умолчанию: 0)
- dampening (float, optional) – ослабление для момента (по умолчанию: 0)
- nesterov (bool, optional) – включает момент Нестерова (по умолчанию: False)
- maximize (bool, optional) – максимизировать параметры на основе целевой функции вместо минимизации (по умолчанию: False)
- foreach (bool, optional) – используется ли реализация foreach оптимизатора. Если пользователь не указал (т.е. foreach равно None), мы попытаемся использовать foreach вместо реализации цикла for на CUDA, так как обычно это значительно производительнее. Обратите внимание, что реализация foreach использует ~ sizeof(params) больше пиковой памяти, чем версия цикла for, из-за того, что промежуточные результаты представляют собой список тензоров, а не просто один тензор. Если память ограничена, передавайте меньше параметров через оптимизатор за раз или установите этот флаг в False (по умолчанию: None)
- differentiable (bool, optional) – должна ли происходить автодифференциация через шаг оптимизатора при обучении. В противном случае функция step() выполняется в контексте torch.no_grad(). Установка в True может ухудшить производительность, поэтому оставьте ее False, если вы не планируете выполнять автодифференциацию через эту экземпляр (по умолчанию: False)
Пример
>>> optimizer = torch.optim.SGD(model.parameters(), lr=0.1, momentum=0.9) >>> optimizer.zero_grad() >>> loss_fn(model(input), target).backward() >>> optimizer.step()
Примечание
Реализация SGD с импульсом/Нестеровым импульсом немного отличается от реализации Sutskever et. al. и реализаций в некоторых других фреймворках.
Рассмотрим конкретный случай импульса, обновление можно записать как
где , , и обозначают параметры, градиент, скорость и импульс соответственно.
Это отличается от Sutskever et. al. и других фреймворков, которые используют обновление в форме
Аналогичным образом модифицируется версия Нестерова.
Кроме того, начальное значение буфера импульса устанавливается в значение градиента на первом шаге. Это отличается от некоторых других фреймворков, которые инициализируют его нулями.
-
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) -
Зарегистрировать обработчик 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 на месте или, по желанию, возвращать новый. Если возвращается 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 на месте или, по желанию, возвращать новый. Зарегистрированный обработчик можно использовать для выполнения пост-обработки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, который будет вызываться перед вызовом
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 -
группа параметров — это словарь. Каждая группа параметров содержит метаданные, специфичные для оптимизатора, такие как скорость обучения и разложение весов, а также список идентификаторов параметров параметров в группе.
-
ПРИМЕЧАНИЕ: Идентификаторы параметров могут выглядеть как индексы, но они являются просто идентификаторами, сопоставляющими состояние с param_group. При загрузке из state_dict, оптимизатор будет связывать param_group
params(целые идентификаторы) и optimizerparam_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.SGD.html