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.Tensors илиdicts. Указывает, какие тензоры следует оптимизировать. - defaults (dict[str, Any]) – (dict): словарь со значениями параметров оптимизации по умолчанию (используются, если группа параметров не задаёт эти значения).
-
params (итерируемый объект) – итерируемый объект из
Добавить группу параметров в | |
Загрузить состояние оптимизатора. | |
Зарегистрировать предварительный перехватчик load_state_dict, который будет вызван перед вызовом | |
Зарегистрировать последующий перехватчик load_state_dict, который будет вызван после вызова | |
Вернуть состояние оптимизатора в виде | |
Зарегистрировать предварительный перехватчик state dict, который будет вызван перед вызовом | |
Зарегистрировать последующий перехватчик state dict, который будет вызван после вызова | |
Выполнить один шаг оптимизации для обновления параметров. | |
Зарегистрировать предварительный перехватчик шага оптимизатора, который будет вызван перед выполнением шага. | |
Зарегистрировать последующий перехватчик шага оптимизатора, который будет вызван после выполнения шага. | |
Обнулить градиенты всех оптимизируемых |
Перехватчики уровня модуля
-
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() в неподходящий момент.
Базовый класс для всех планировщиков скорости обучения. | |
Устанавливает начальную скорость обучения. | |
Умножает скорость обучения каждой группы параметров на коэффициент, заданный указанной функцией. | |
Уменьшает скорость обучения каждой группы параметров в gamma раз каждые step_size эпох. | |
Уменьшает скорость обучения каждой группы параметров в gamma раз, когда номер эпохи достигает одной из контрольных точек. | |
Умножает скорость обучения каждой группы параметров на небольшой постоянный коэффициент. | |
Уменьшает скорость обучения каждой группы параметров, линейно изменяя небольшой множительный коэффициент. | |
Уменьшает скорость обучения каждой группы параметров в gamma раз каждую эпоху. | |
Уменьшает скорость обучения каждой группы параметров с помощью полиномиальной функции за заданное число итераций total_iters. | |
Устанавливает скорость обучения каждой группы параметров в соответствии с расписанием косинусного отжига. | |
Объединяет список планировщиков скорости обучения в цепочку. | |
Содержит список планировщиков, которые должны вызываться последовательно в процессе оптимизации. | |
Уменьшает скорость обучения, когда улучшение метрики прекращается. | |
Устанавливает скорость обучения каждой группы параметров в соответствии с политикой циклической скорости обучения (CLR). | |
Устанавливает скорость обучения каждой группы параметров в соответствии с политикой скорости обучения 1cycle. | |
Устанавливает скорость обучения каждой группы параметров в соответствии с расписанием косинусного отжига. |
Как использовать именованные параметры для загрузки словаря состояний оптимизатора
Функция 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:
где 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) и экспоненциального скользящего среднего (EMA). | |
Снижает скорость обучения в каждой группе параметров до фиксированного значения. | |
Возвращает функцию применения экспоненциального скользящего среднего (EMA) к нескольким параметрам. | |
Возвращает функцию применения стохастического усреднения весов (SWA) к одному параметру. | |
Возвращает функцию применения стохастического усреднения весов (SWA) к нескольким параметрам. |
-
torch.optim.swa_utils.get_ema_multi_avg_fn(decay=0.999)[исходный код] -
Возвращает функцию применения экспоненциального скользящего среднего (EMA) к нескольким параметрам.
EMA вычисляется следующим образом:
где — параметр EMA на шаге , — параметр модели на шаге , а — скорость затухания (по умолчанию: 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