Spec-Zone.ru › PyTorch 2.14

torch.optim

Создано: Jun 13, 2025 | Последнее обновление: May 10, 2026

torch.optim — это пакет, реализующий различные алгоритмы оптимизации.

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

Использование оптимизатора

Чтобы использовать torch.optim, необходимо создать объект оптимизатора, который будет хранить текущее состояние и обновлять параметры на основе вычисленных градиентов.

Создание

Чтобы создать Optimizer, необходимо передать ему итерируемый объект, содержащий параметры (все они должны быть Parameter s), или именованные параметры (кортежи вида (str, Parameter)), которые нужно оптимизировать. Затем можно указать параметры, специфичные для оптимизатора, например скорость обучения, затухание весов и т. д.

Пример:

optimizer = optim.SGD(model.parameters(), lr=0.01, momentum=0.9)
optimizer = optim.Adam([var1, var2], lr=0.0001)

Пример с именованными параметрами:

optimizer = optim.SGD(model.named_parameters(), lr=0.01, momentum=0.9)
optimizer = optim.Adam([('layer0', var1), ('layer1', var2)], lr=0.0001)

Параметры для отдельных параметров

Optimizer s также поддерживают указание параметров для отдельных параметров. Для этого вместо передачи итерируемого объекта из Variable s передайте итерируемый объект из dict s. Каждый из них определяет отдельную группу параметров и должен содержать ключ params со списком принадлежащих этой группе параметров. Остальные ключи должны соответствовать именованным аргументам, принимаемым оптимизаторами, и будут использоваться как параметры оптимизации для этой группы.

Например, это очень полезно, когда нужно задать скорость обучения для каждого слоя:

optim.SGD([
    {'params': model.base.parameters(), 'lr': 1e-2},
    {'params': model.classifier.parameters()}
], lr=1e-3, momentum=0.9)

optim.SGD([
    {'params': model.base.named_parameters(), 'lr': 1e-2},
    {'params': model.classifier.named_parameters()}
], lr=1e-3, momentum=0.9)

Это означает, что для параметров model.base будет использоваться скорость обучения 1e-2, тогда как для параметров model.classifier сохранится скорость обучения по умолчанию 1e-3. Наконец, для всех параметров будет использоваться моментум 0.9.

Примечание

Параметры по-прежнему можно передавать в виде именованных аргументов. Они будут использоваться как значения по умолчанию для групп, в которых эти параметры не переопределены. Это удобно, если нужно изменить только один параметр, сохранив все остальные одинаковыми для всех групп параметров.

Также рассмотрите следующий пример, в котором параметрам назначаются разные штрафы. Помните, что parameters() возвращает итерируемый объект, содержащий все обучаемые параметры, в том числе смещения и другие параметры, для которых может потребоваться отдельный штраф. Для этого можно задать отдельные веса штрафа для каждой группы параметров:

bias_params = [p for name, p in self.named_parameters() if 'bias' in name]
others = [p for name, p in self.named_parameters() if 'bias' not in name]

optim.SGD([
    {'params': others},
    {'params': bias_params, 'weight_decay': 0}
], weight_decay=1e-2, lr=1e-2)

Таким образом, параметры смещения отделяются от остальных параметров, а для параметров смещения отдельно устанавливается weight_decay, равный 0, чтобы не применять к этой группе никаких штрафов.

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

Все оптимизаторы реализуют метод step(), который обновляет параметры. Его можно использовать двумя способами:

optimizer.step()

Это упрощённый вариант, поддерживаемый большинством оптимизаторов. Функцию можно вызвать после вычисления градиентов, например с помощью backward().

Пример:

for input, target in dataset:
    optimizer.zero_grad()
    output = model(input)
    loss = loss_fn(output, target)
    loss.backward()
    optimizer.step()

optimizer.step(closure)

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

Пример:

for input, target in dataset:
    def closure():
        optimizer.zero_grad()
        output = model(input)
        loss = loss_fn(output, target)
        loss.backward()
        return loss
    optimizer.step(closure)

Базовый класс

class torch.optim.Optimizer(params, defaults) [исходный код]

Базовый класс для всех оптимизаторов.

Предупреждение

Параметры необходимо задавать в виде коллекций с детерминированным порядком, который не меняется при повторных запусках. Примеры объектов, не удовлетворяющих этим требованиям, — множества и итераторы по значениям словарей.

Параметры:
  • params (итерируемый объект) – итерируемый объект из torch.Tensor s или dict s. Указывает, какие тензоры следует оптимизировать.
  • defaults (dict[str, Any]) – (dict): словарь со значениями параметров оптимизации по умолчанию (используются, если группа параметров не задаёт эти значения).

Optimizer.add_param_group

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

Optimizer.load_state_dict

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

Optimizer.register_load_state_dict_pre_hook

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

Optimizer.register_load_state_dict_post_hook

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

Optimizer.state_dict

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

Optimizer.register_state_dict_pre_hook

Зарегистрировать предварительный перехватчик state dict, который будет вызван перед вызовом state_dict().

Optimizer.register_state_dict_post_hook

Зарегистрировать последующий перехватчик state dict, который будет вызван после вызова state_dict().

Optimizer.step

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

Optimizer.register_step_pre_hook

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

Optimizer.register_step_post_hook

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

Optimizer.zero_grad

Обнулить градиенты всех оптимизируемых torch.Tensor s.

Перехватчики уровня модуля

torch.optim.optimizer.register_optimizer_step_post_hook(hook) [исходный код]

Зарегистрировать последующий перехватчик, общий для всех оптимизаторов.

Перехватчик должен иметь следующую сигнатуру:

hook(optimizer, args, kwargs) -> None
Параметры:

hook (Callable) – Пользовательский перехватчик, регистрируемый для всех оптимизаторов.

Возвращает:

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

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

torch.utils.hooks.RemovableHandle

torch.optim.optimizer.register_optimizer_step_pre_hook(hook) [исходный код]

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

Перехватчик должен иметь следующую сигнатуру:

hook(optimizer, args, kwargs) -> None or modified args and kwargs
Параметры:

hook (Callable) – Пользовательский перехватчик, регистрируемый для всех оптимизаторов.

Возвращает:

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

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

torch.utils.hooks.RemovableHandle

Утилиты

torch.optim.swap_in_optimizer_params_and_state(optimizer, swapin_parameters, swapin_optim_state) [исходный код]

Временно заменить параметры и состояние оптимизатора переданными параметрами и состоянием оптимизатора, а при выходе восстановить исходные значения.

На время работы контекстного менеджера все API оптимизатора (включая пользовательские перехватчики) используют подставленные значения; при выходе исходное состояние оптимизатора восстанавливается.

Отличие этого API от optimizer.load_state_dict заключается в том, что optimizer.load_state_dict обновляет только состояние оптимизатора, не изменяя параметры в param_groups. Этот API также подставляет параметры, поэтому optimizer.step() выполняется над переданными вами подставленными тензорами параметров.

Параметры:
  • optimizer (Optimizer) – действующий оптимизатор; его состояние должно быть уже инициализировано.
  • swapin_parameters (dict[str, Tensor]) – тензоры, используемые в качестве параметров на время работы контекстного менеджера; они должны быть указаны в том же порядке, что и существующие входные параметры оптимизатора (чаще всего в порядке model.named_parameters()).
  • swapin_optim_state (dict[str, Any]) – словарь формы optimizer.state_dict() ({"state": ..., "param_groups": ...}), содержащий устанавливаемое состояние. "state" индексируется упакованными целочисленными идентификаторами параметров, а "param_groups" повторяет структуру optimizer.param_groups: каждая запись "params" представляет собой список таких упакованных идентификаторов, а остальные ключи содержат гиперпараметры группы (lr, betas, foreach, capturable, …). Только изменения тензоров на месте передаются в предоставленный пользователем объект swapin_optim_state; все прочие побочные эффекты (например, присваивание нового тензора состоянию оптимизатора) игнорируются.

Пример

Этот API можно использовать, чтобы запускать optimizer.step() на FakeTensor версиях параметров и состояния для трассировки без строгих ограничений — захватывая граф FX шага, не затрагивая действующий оптимизатор:

from torch.fx.experimental.proxy_tensor import make_fx
from torch._subclasses import FakeTensorMode
from torch.utils import _pytree as pytree

fake_mode = FakeTensorMode(allow_non_fake_inputs=True)
with fake_mode:
    fake_params = {
        n: fake_mode.from_tensor(p) for n, p in model.named_parameters()
    }
    fake_osd = pytree.tree_map_only(
        torch.Tensor, fake_mode.from_tensor, optimizer.state_dict()
    )


def step_fn(params, osd):
    with swap_in_optimizer_params_and_state(optimizer, params, osd):
        optimizer.step()
    return params, osd


gm = make_fx(step_fn)(fake_params, fake_osd)

Алгоритмы

Adadelta

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

Adafactor

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

Adagrad

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

Adam

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

AdamW

Реализует алгоритм AdamW, в котором затухание весов не накапливается ни в моментуме, ни в дисперсии.

SparseAdam

SparseAdam реализует версию алгоритма Adam с маскированием, подходящую для разреженных градиентов.

Adamax

Реализует алгоритм Adamax (вариант Adam, основанный на бесконечной норме).

ASGD

Реализует усреднённый стохастический градиентный спуск.

LBFGS

Реализует алгоритм L-BFGS.

Muon

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

NAdam

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

RAdam

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

RMSprop

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

Rprop

Реализует алгоритм устойчивого обратного распространения ошибки.

SGD

Реализует стохастический градиентный спуск (при необходимости с моментумом).

Для многих наших алгоритмов доступны различные реализации, оптимизированные по производительности, удобочитаемости и/или универсальности. Поэтому, если пользователь не указал конкретную реализацию, по умолчанию мы выбираем в целом самую быструю для текущего устройства.

У нас есть три основные категории реализаций: циклы for, foreach (несколько тензоров) и fused. Самые простые реализации выполняют большие блоки вычислений в цикле for по параметрам. Циклы for обычно медленнее наших реализаций foreach, которые объединяют параметры в несколько тензоров и выполняют большие блоки вычислений одновременно, сокращая число последовательных вызовов ядер. Некоторые оптимизаторы имеют ещё более быстрые объединённые реализации, которые сливают большие блоки вычислений в одно ядро. Реализации foreach можно считать горизонтальным слиянием, а fused — вертикальным слиянием поверх него.

В целом производительность трёх реализаций располагается в таком порядке: fused > foreach > for-loop. Поэтому, когда это возможно, мы по умолчанию выбираем foreach вместо for-loop. Это возможно, если доступна реализация foreach, пользователь не указал аргументы, специфичные для какой-либо реализации (например, fused, foreach, differentiable), а все тензоры являются нативными. Обратите внимание: хотя fused должна работать ещё быстрее, чем foreach, эти реализации появились недавно, и мы хотим дать им больше времени на проверку, прежде чем повсеместно переключиться на них. В таблице ниже приведён статус стабильности каждой реализации; при этом вы можете попробовать их уже сейчас!

В таблице ниже показаны доступные реализации каждого алгоритма и используемые по умолчанию:

Алгоритм

По умолчанию

Есть foreach?

Есть fused?

Adadelta

foreach

да

нет

Adafactor

for-loop

нет

нет

Adagrad

foreach

да

да (только CPU)

Adam

foreach

да

да

AdamW

foreach

да

да

SparseAdam

for-loop

нет

нет

Adamax

foreach

да

нет

ASGD

foreach

да

нет

LBFGS

for-loop

нет

нет

Muon

for-loop

нет

нет

NAdam

foreach

да

нет

RAdam

foreach

да

нет

RMSprop

foreach

да

нет

Rprop

foreach

да

нет

SGD

foreach

да

да

В таблице ниже показан статус стабильности объединённых реализаций:

Алгоритм

CPU

CUDA

MPS

Adadelta

не поддерживается

не поддерживается

не поддерживается

Adafactor

не поддерживается

не поддерживается

не поддерживается

Adagrad

бета

не поддерживается

не поддерживается

Adam

бета

стабильная

бета

AdamW

бета

стабильная

бета

SparseAdam

не поддерживается

не поддерживается

не поддерживается

Adamax

не поддерживается

не поддерживается

не поддерживается

ASGD

не поддерживается

не поддерживается

не поддерживается

LBFGS

не поддерживается

не поддерживается

не поддерживается

Muon

не поддерживается

не поддерживается

не поддерживается

NAdam

не поддерживается

не поддерживается

не поддерживается

RAdam

не поддерживается

не поддерживается

не поддерживается

RMSprop

не поддерживается

не поддерживается

не поддерживается

Rprop

не поддерживается

не поддерживается

не поддерживается

SGD

бета

бета

бета

Как настроить скорость обучения

torch.optim.lr_scheduler.LRScheduler предоставляет несколько методов настройки скорости обучения в зависимости от количества эпох. torch.optim.lr_scheduler.ReduceLROnPlateau позволяет динамически снижать скорость обучения на основе некоторых показателей валидации.

Планировщик скорости обучения следует применять после обновления оптимизатора; например, код следует написать так:

Пример:

optimizer = optim.SGD(model.parameters(), lr=0.01, momentum=0.9)
scheduler = ExponentialLR(optimizer, gamma=0.9)

for epoch in range(20):
    for input, target in dataset:
        optimizer.zero_grad()
        output = model(input)
        loss = loss_fn(output, target)
        loss.backward()
        optimizer.step()
    scheduler.step()

Большинство планировщиков скорости обучения можно вызывать один за другим (это также называется объединением планировщиков в цепочку). В результате каждый планировщик последовательно применяется к скорости обучения, полученной от предыдущего.

Пример:

optimizer = optim.SGD(model.parameters(), lr=0.01, momentum=0.9)
scheduler1 = ExponentialLR(optimizer, gamma=0.9)
scheduler2 = MultiStepLR(optimizer, milestones=[30,80], gamma=0.1)

for epoch in range(20):
    for input, target in dataset:
        optimizer.zero_grad()
        output = model(input)
        loss = loss_fn(output, target)
        loss.backward()
        optimizer.step()
    scheduler1.step()
    scheduler2.step()

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

>>> scheduler = ...
>>> for epoch in range(100):
>>>     train(...)
>>>     validate(...)
>>>     scheduler.step()

Предупреждение

До PyTorch 1.1.0 планировщик скорости обучения следовало вызывать до обновления оптимизатора; в версии 1.1.0 это поведение изменилось, что нарушило обратную совместимость. Если вы используете планировщик скорости обучения (вызывая scheduler.step()) до обновления оптимизатора (вызывая optimizer.step()), первое значение расписания скорости обучения будет пропущено. Если после обновления до PyTorch 1.1.0 вам не удаётся воспроизвести результаты, проверьте, не вызываете ли вы scheduler.step() в неподходящий момент.

lr_scheduler.LRScheduler

Базовый класс для всех планировщиков скорости обучения.

lr_scheduler.LambdaLR

Устанавливает начальную скорость обучения.

lr_scheduler.MultiplicativeLR

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

lr_scheduler.StepLR

Уменьшает скорость обучения каждой группы параметров в gamma раз каждые step_size эпох.

lr_scheduler.MultiStepLR

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

lr_scheduler.ConstantLR

Умножает скорость обучения каждой группы параметров на небольшой постоянный коэффициент.

lr_scheduler.LinearLR

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

lr_scheduler.ExponentialLR

Уменьшает скорость обучения каждой группы параметров в gamma раз каждую эпоху.

lr_scheduler.PolynomialLR

Уменьшает скорость обучения каждой группы параметров с помощью полиномиальной функции за заданное число итераций total_iters.

lr_scheduler.CosineAnnealingLR

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

lr_scheduler.ChainedScheduler

Объединяет список планировщиков скорости обучения в цепочку.

lr_scheduler.SequentialLR

Содержит список планировщиков, которые должны вызываться последовательно в процессе оптимизации.

lr_scheduler.ReduceLROnPlateau

Уменьшает скорость обучения, когда улучшение метрики прекращается.

lr_scheduler.CyclicLR

Устанавливает скорость обучения каждой группы параметров в соответствии с политикой циклической скорости обучения (CLR).

lr_scheduler.OneCycleLR

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

lr_scheduler.CosineAnnealingWarmRestarts

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

Как использовать именованные параметры для загрузки словаря состояний оптимизатора

Функция load_state_dict() сохраняет необязательное содержимое param_names из загруженного словаря состояний, если оно присутствует. Однако процесс загрузки состояния оптимизатора не меняется, поскольку для обеспечения совместимости важен порядок параметров (на случай различий в порядке). Чтобы использовать имена параметров из загруженного словаря состояний, необходимо реализовать пользовательский register_load_state_dict_pre_hook в соответствии с требуемым поведением.

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

Пример:

class OneLayerModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc = nn.Linear(3, 4)

    def forward(self, x):
        return self.fc(x)

model = OneLayerModel()
optimizer = optim.SGD(model.named_parameters(), lr=0.01, momentum=0.9)
# training..
torch.save(optimizer.state_dict(), PATH)

Предположим, что model реализует экспертную модель (MoE), и мы хотим продублировать её и возобновить обучение для двух экспертов, инициализированных так же, как слой fc. Для следующего model2 мы создаём два слоя, идентичных fc, и возобновляем обучение, загружая веса модели и состояния оптимизатора из model в оба fc1 и fc2 объекта model2 (и соответствующим образом корректируя их):

class TwoLayerModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(3, 4)
        self.fc2 = nn.Linear(3, 4)

    def forward(self, x):
        return (self.fc1(x) + self.fc2(x)) / 2

model2 = TwoLayerModel()
# adapt and load model weights..
optimizer2 = optim.SGD(model2.named_parameters(), lr=0.01, momentum=0.9)

Чтобы загрузить словарь состояний для optimizer2 со словарём состояний предыдущего оптимизатора так, чтобы и fc1, и fc2 были инициализированы копией состояний оптимизатора fc (для возобновления обучения каждого слоя с fc), можно использовать следующий хук:

def adapt_state_dict_ids(optimizer, state_dict):
    adapted_state_dict = deepcopy(optimizer.state_dict())
    # Copy setup parameters (lr, weight_decay, etc.), in case they differ in the loaded state dict.
    for k, v in state_dict['param_groups'][0].items():
        if k not in ['params', 'param_names']:
            adapted_state_dict['param_groups'][0][k] = v

    lookup_dict = {
        'fc1.weight': 'fc.weight',
        'fc1.bias': 'fc.bias',
        'fc2.weight': 'fc.weight',
        'fc2.bias': 'fc.bias'
    }
    clone_deepcopy = lambda d: {k: (v.clone() if isinstance(v, torch.Tensor) else deepcopy(v)) for k, v in d.items()}
    for param_id, param_name in zip(
            optimizer.state_dict()['param_groups'][0]['params'],
            optimizer.state_dict()['param_groups'][0]['param_names']):
        name_in_loaded = lookup_dict[param_name]
        index_in_loaded_list = state_dict['param_groups'][0]['param_names'].index(name_in_loaded)
        id_in_loaded = state_dict['param_groups'][0]['params'][index_in_loaded_list]
        # Copy the state of the corresponding parameter
        if id_in_loaded in state_dict['state']:
            adapted_state_dict['state'][param_id] = clone_deepcopy(state_dict['state'][id_in_loaded])

    return adapted_state_dict

optimizer2.register_load_state_dict_pre_hook(adapt_state_dict_ids)
optimizer2.load_state_dict(torch.load(PATH)) # The previous optimizer saved state_dict

Это гарантирует, что при загрузке модели будет использоваться адаптированный state_dict с правильными состояниями слоёв model2. Обратите внимание, что этот код предназначен специально для данного примера (например, предполагается наличие одной группы параметров); в других случаях могут потребоваться иные настройки.

В следующем примере показано, как обработать отсутствующие параметры в загруженном state dict при изменении структуры модели. В Model_bypass добавляется новый слой bypass, которого нет в исходной Model1. Для возобновления обучения используется пользовательский хук adapt_state_dict_missing_param, адаптирующий state_dict оптимизатора и обеспечивающий правильное сопоставление существующих параметров, в то время как отсутствующие параметры (например, слой обхода) остаются без изменений (как инициализировано в этом примере). Такой подход позволяет без проблем загружать состояние оптимизатора и возобновлять обучение, несмотря на изменения модели. Новый слой обхода будет обучаться с нуля:

class Model1(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc = nn.Linear(5, 5)

    def forward(self, x):
        return self.fc(x) + x


model = Model1()
optimizer = optim.SGD(model.named_parameters(), lr=0.01, momentum=0.9)
# training..
torch.save(optimizer.state_dict(), PATH)

class Model_bypass(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc = nn.Linear(5, 5)
        self.bypass = nn.Linear(5, 5, bias=False)
        torch.nn.init.eye_(self.bypass.weight)

    def forward(self, x):
        return self.fc(x) + self.bypass(x)

model2 = Model_bypass()
optimizer2 = optim.SGD(model2.named_parameters(), lr=0.01, momentum=0.9)

def adapt_state_dict_missing_param(optimizer, state_dict):
    adapted_state_dict = deepcopy(optimizer.state_dict())
    # Copy setup parameters (lr, weight_decay, etc.), in case they differ in the loaded state dict.
    for k, v in state_dict['param_groups'][0].items():
        if k not in ['params', 'param_names']:
            adapted_state_dict['param_groups'][0][k] = v

    lookup_dict = {
        'fc.weight': 'fc.weight',
        'fc.bias': 'fc.bias',
        'bypass.weight': None,
    }

    clone_deepcopy = lambda d: {k: (v.clone() if isinstance(v, torch.Tensor) else deepcopy(v)) for k, v in d.items()}
    for param_id, param_name in zip(
            optimizer.state_dict()['param_groups'][0]['params'],
            optimizer.state_dict()['param_groups'][0]['param_names']):
        name_in_loaded = lookup_dict[param_name]
        if name_in_loaded in state_dict['param_groups'][0]['param_names']:
            index_in_loaded_list = state_dict['param_groups'][0]['param_names'].index(name_in_loaded)
            id_in_loaded = state_dict['param_groups'][0]['params'][index_in_loaded_list]
            # Copy the state of the corresponding parameter
            if id_in_loaded in state_dict['state']:
                adapted_state_dict['state'][param_id] = clone_deepcopy(state_dict['state'][id_in_loaded])

    return adapted_state_dict

optimizer2.register_load_state_dict_pre_hook(adapt_state_dict_ids)
optimizer2.load_state_dict(torch.load(PATH)) # The previous optimizer saved state_dict

В третьем примере вместо загрузки состояния в соответствии с порядком параметров (по умолчанию) этот хук используется для загрузки по именам параметров:

def names_matching(optimizer, state_dict):
    assert len(state_dict['param_groups']) == len(optimizer.state_dict()['param_groups'])
    adapted_state_dict = deepcopy(optimizer.state_dict())
    for g_ind in range(len(state_dict['param_groups'])):
        assert len(state_dict['param_groups'][g_ind]['params']) == len(
            optimizer.state_dict()['param_groups'][g_ind]['params'])

        for k, v in state_dict['param_groups'][g_ind].items():
            if k not in ['params', 'param_names']:
                adapted_state_dict['param_groups'][g_ind][k] = v

        for param_id, param_name in zip(
                optimizer.state_dict()['param_groups'][g_ind]['params'],
                optimizer.state_dict()['param_groups'][g_ind]['param_names']):
            index_in_loaded_list = state_dict['param_groups'][g_ind]['param_names'].index(param_name)
            id_in_loaded = state_dict['param_groups'][g_ind]['params'][index_in_loaded_list]
            # Copy the state of the corresponding parameter
            if id_in_loaded in state_dict['state']:
                adapted_state_dict['state'][param_id] = deepcopy(state_dict['state'][id_in_loaded])

    return adapted_state_dict

Усреднение весов (SWA и EMA)

torch.optim.swa_utils.AveragedModel реализует стохастическое усреднение весов (SWA) и экспоненциальное скользящее среднее (EMA), torch.optim.swa_utils.SWALR реализует планировщик скорости обучения SWA, а torch.optim.swa_utils.update_bn() — это вспомогательная функция для обновления статистик пакетной нормализации SWA/EMA в конце обучения.

Метод SWA предложен в статье Усреднение весов приводит к более широким оптимумам и лучшей обобщающей способности.

EMA — широко известный метод сокращения времени обучения за счёт уменьшения числа необходимых обновлений весов. Это разновидность усреднения Поляка, но вместо равных весов на всех итерациях используются экспоненциальные веса.

Создание усреднённых моделей

Класс AveragedModel предназначен для вычисления весов модели SWA или EMA.

Создать усреднённую модель SWA можно следующим образом:

>>> averaged_model = AveragedModel(model)

Модели EMA создаются путём указания аргумента multi_avg_fn следующим образом:

>>> decay = 0.999
>>> averaged_model = AveragedModel(model, multi_avg_fn=get_ema_multi_avg_fn(decay))

Затухание — это параметр от 0 до 1, определяющий скорость затухания усреднённых параметров. Если не передать его в torch.optim.swa_utils.get_ema_multi_avg_fn(), значение по умолчанию будет равно 0.999. Значение затухания должно быть близко к 1.0, поскольку меньшие значения могут вызвать проблемы со сходимостью оптимизации.

torch.optim.swa_utils.get_ema_multi_avg_fn() возвращает функцию, применяющую к весам следующее уравнение EMA:

W0EMA=W0modelW_0^{\text{EMA}} = W_0^{\text{model}}
Wt+1EMA=decay×WtEMA+(1−decay)×Wt+1modelW_{t+1}^{\text{EMA}} = \text{decay} \times W_t^{\text{EMA}} + (1 - \text{decay}) \times W_{t+1}^{\text{model}}

где W_t^{\text{EMA}} — параметр EMA на шаге t, W_t^{\text{model}} — параметр модели на шаге t, а decay — скорость затухания EMA (по умолчанию: 0.999).

Здесь модель model может быть произвольным объектом torch.nn.Module. averaged_model будет отслеживать скользящие средние параметров model. Для обновления этих средних следует использовать функцию update_parameters() после optimizer.step():

>>> averaged_model.update_parameters(model)

При использовании SWA и EMA этот вызов обычно выполняется сразу после step() оптимизатора. В случае SWA его обычно пропускают в течение некоторого количества начальных шагов обучения.

Пользовательские стратегии усреднения

По умолчанию torch.optim.swa_utils.AveragedModel вычисляет текущее простое среднее предоставленных параметров, но также можно использовать пользовательские функции усреднения с параметрами avg_fn или multi_avg_fn:

  • avg_fn позволяет задать функцию, обрабатывающую каждый кортеж параметров (усреднённый параметр, параметр модели) и возвращающую новый усреднённый параметр.
  • multi_avg_fn позволяет задавать более эффективные операции, одновременно обрабатывающие кортеж списков параметров (список усреднённых параметров, список параметров модели), например, с помощью функций torch._foreach*. Эта функция должна обновлять усреднённые параметры на месте.

В следующем примере ema_model вычисляет экспоненциальное скользящее среднее с помощью параметра avg_fn:

>>> ema_avg = lambda averaged_model_parameter, model_parameter, num_averaged:\
>>>         0.9 * averaged_model_parameter + 0.1 * model_parameter
>>> ema_model = torch.optim.swa_utils.AveragedModel(model, avg_fn=ema_avg)

В следующем примере ema_model вычисляет экспоненциальное скользящее среднее с помощью более эффективного параметра multi_avg_fn:

>>> ema_model = AveragedModel(model, multi_avg_fn=get_ema_multi_avg_fn(0.9))

Расписания скорости обучения SWA

Как правило, при использовании SWA скорость обучения устанавливается на постоянное высокое значение. SWALR — это планировщик скорости обучения, который снижает её до фиксированного значения, а затем поддерживает на этом уровне. Например, следующий код создаёт планировщик, который линейно снижает скорость обучения с начального значения до 0.05 за 5 эпох в каждой группе параметров:

>>> swa_scheduler = torch.optim.swa_utils.SWALR(optimizer, \
>>>         anneal_strategy="linear", anneal_epochs=5, swa_lr=0.05)

Вместо линейного снижения до фиксированного значения можно использовать косинусное снижение, задав anneal_strategy="cos".

Обработка пакетной нормализации

update_bn() — это вспомогательная функция, которая позволяет вычислить статистики пакетной нормализации для модели SWA на заданном загрузчике данных loader в конце обучения:

>>> torch.optim.swa_utils.update_bn(loader, swa_model)

update_bn() применяет swa_model к каждому элементу загрузчика данных и вычисляет статистики активаций для каждого слоя пакетной нормализации модели.

Предупреждение

update_bn() предполагает, что каждый пакет данных в загрузчике loader представляет собой либо тензор, либо список тензоров, где первый элемент — тензор, к которому следует применить сеть swa_model. Если загрузчик данных имеет другую структуру, статистики пакетной нормализации для swa_model можно обновить, выполнив прямой проход с помощью swa_model для каждого элемента набора данных.

Всё вместе: SWA

В приведённом ниже примере swa_model — это модель SWA, которая накапливает средние значения весов. Мы обучаем модель в течение 300 эпох и переключаемся на расписание скорости обучения SWA и начинаем собирать средние значения параметров SWA на эпохе 160:

>>> loader, optimizer, model, loss_fn = ...
>>> swa_model = torch.optim.swa_utils.AveragedModel(model)
>>> scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=300)
>>> swa_start = 160
>>> swa_scheduler = SWALR(optimizer, swa_lr=0.05)
>>>
>>> for epoch in range(300):
>>>       for input, target in loader:
>>>           optimizer.zero_grad()
>>>           loss_fn(model(input), target).backward()
>>>           optimizer.step()
>>>       if epoch > swa_start:
>>>           swa_model.update_parameters(model)
>>>           swa_scheduler.step()
>>>       else:
>>>           scheduler.step()
>>>
>>> # Update bn statistics for the swa_model at the end
>>> torch.optim.swa_utils.update_bn(loader, swa_model)
>>> # Use swa_model to make predictions on test data
>>> preds = swa_model(test_input)

Всё вместе: EMA

В приведённом ниже примере ema_model — это модель EMA, накапливающая экспоненциально затухающие средние значения весов со скоростью затухания 0.999. Мы обучаем модель в течение 300 эпох и сразу начинаем собирать средние значения EMA.

>>> loader, optimizer, model, loss_fn = ...
>>> ema_model = torch.optim.swa_utils.AveragedModel(model, \
>>>             multi_avg_fn=torch.optim.swa_utils.get_ema_multi_avg_fn(0.999))
>>>
>>> for epoch in range(300):
>>>       for input, target in loader:
>>>           optimizer.zero_grad()
>>>           loss_fn(model(input), target).backward()
>>>           optimizer.step()
>>>           ema_model.update_parameters(model)
>>>
>>> # Update bn statistics for the ema_model at the end
>>> torch.optim.swa_utils.update_bn(loader, ema_model)
>>> # Use ema_model to make predictions on test data
>>> preds = ema_model(test_input)

swa_utils.AveragedModel

Реализует усреднённую модель для стохастического усреднения весов (SWA) и экспоненциального скользящего среднего (EMA).

swa_utils.SWALR

Снижает скорость обучения в каждой группе параметров до фиксированного значения.

swa_utils.get_ema_avg_fn

Возвращает функцию применения экспоненциального скользящего среднего (EMA) к нескольким параметрам.

swa_utils.get_swa_avg_fn

Возвращает функцию применения стохастического усреднения весов (SWA) к одному параметру.

swa_utils.get_swa_multi_avg_fn

Возвращает функцию применения стохастического усреднения весов (SWA) к нескольким параметрам.

torch.optim.swa_utils.get_ema_multi_avg_fn(decay=0.999) [исходный код]

Возвращает функцию применения экспоненциального скользящего среднего (EMA) к нескольким параметрам.

EMA вычисляется следующим образом:

W0EMA=W0modelW_0^{\text{EMA}} = W_0^{\text{model}}
Wt+1EMA=decay×WtEMA+(1−decay)×Wt+1modelW_{t+1}^{\text{EMA}} = \text{decay} \times W_t^{\text{EMA}} + (1 - \text{decay}) \times W_{t+1}^{\text{model}}

где WtEMAW_t^{\text{EMA}} — параметр EMA на шаге tt, WtmodelW_t^{\text{model}} — параметр модели на шаге tt, а decay\text{decay} — скорость затухания (по умолчанию: 0.999).

Параметры:

decay (float) – Скорость затухания EMA. Должна находиться в диапазоне [0, 1]. По умолчанию: 0.999

Возвращает:

Функция, обновляющая параметры EMA с учётом текущих параметров модели

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

Callable

torch.optim.swa_utils.update_bn(loader, model, device=None) [исходный код]

Обновляет буферы running_mean и running_var модели BatchNorm.

Выполняет один проход по данным в loader для оценки статистик активаций слоёв BatchNorm модели.

Параметры:
  • loader (torch.utils.data.DataLoader) – загрузчик набора данных, используемый для вычисления статистик активаций. Каждый пакет данных должен представлять собой либо тензор, либо список/кортеж, первым элементом которого является тензор с данными.
  • model (torch.nn.Module) – модель, для которой требуется обновить статистики BatchNorm.
  • device (torch.device, optional) – если задано, данные будут переданы на device перед передачей в model.

Пример

>>> loader, model = ...
>>> torch.optim.swa_utils.update_bn(loader, model)

Примечание

Вспомогательная функция update_bn предполагает, что каждый пакет данных в loader представляет собой либо тензор, либо список или кортеж тензоров; в последнем случае предполагается, что model.forward() следует вызывать для первого элемента списка или кортежа, соответствующего пакету данных.

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/optim.html

Spec-Zone.ru

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