SparseAdam
-
class torch.optim.SparseAdam(params, lr=0.001, betas=(0.9, 0.999), eps=1e-08, maximize=False)[исходный код] -
SparseAdam реализует версию алгоритма Adam с маскированием, подходящую для разреженных градиентов. В настоящее время из-за ограничений реализации (описанных ниже) SparseAdam предназначен только для узкого набора сценариев использования, а именно для параметров с плотным форматом хранения и градиентами с разреженным форматом хранения. Это происходит в особом случае, когда при обратном проходе модуля градиенты уже создаются в разреженном формате хранения. Примером модуля нейронной сети с таким поведением является
nn.Embedding(sparse=True).SparseAdam приближает алгоритм Adam, маскируя обновления параметров и моментов, соответствующие нулевым значениям градиентов. В то время как алгоритм Adam обновляет первый момент, второй момент и параметры на основе всех значений градиентов, SparseAdam обновляет только моменты и параметры, соответствующие ненулевым значениям градиентов.
Упрощенно реализацию
intendedможно представить следующим образом:- Создать маску ненулевых значений разреженных градиентов. Например, если градиент выглядит как [0, 5, 0, 0, 9], маска будет выглядеть как [0, 1, 0, 0, 1].
- Применить эту маску к накапливаемым моментам и выполнить вычисления только для ненулевых значений.
- Применить эту маску к параметрам и обновлять только ненулевые значения.
На практике для оптимизации этого приближения мы используем тензоры с разреженным форматом хранения, а это означает, что чем больше градиентов маскируется за счет отсутствия их материализации, тем эффективнее оптимизация. Поскольку мы полагаемся на использование тензоров с разреженным форматом хранения, мы считаем, что любое материализованное значение в разреженном формате хранения ненулевое, и НЕ проверяем фактически, что все значения ненулевые! Важно не смешивать семантически разреженный тензор (тензор, многие значения которого равны нулю) с тензором с разреженным форматом хранения (тензором, для которого
.is_sparseвозвращаетTrue). Приближение SparseAdam предназначено дляsemanticallyразреженных тензоров, а разреженный формат хранения является лишь деталью реализации. Более наглядная реализация использовала бы MaskedTensors, но они являются экспериментальными.Примечание
Если вы считаете, что ваши градиенты семантически разреженные (но не имеют разреженного формата хранения), этот вариант может вам не подойти. В идеале следует избегать материализации всего, что предположительно является разреженным, поскольку необходимость преобразовать все градиенты из плотного формата хранения в разреженный может свести на нет выигрыш в производительности. В этом случае Adam может быть лучшей альтернативой, если только вы не можете легко настроить модуль так, чтобы он выдавал разреженные градиенты, подобные
nn.Embedding(sparse=True). Если вы настаиваете на преобразовании градиентов, можно вручную переопределить поля.gradпараметров их разреженными эквивалентами перед вызовом.step().- Параметры:
-
- params (iterable) – итерируемый объект с параметрами или именованными параметрами для оптимизации либо итерируемый объект словарей, определяющих группы параметров. При использовании именованных параметров все параметры во всех группах должны иметь имена
- lr (float, Tensor, необязательный) – скорость обучения (по умолчанию: 1e-3)
- betas (Tuple[float, float], необязательный) – коэффициенты для вычисления скользящих средних градиента и его квадрата (по умолчанию: (0.9, 0.999))
- eps (float, необязательный) – слагаемое, добавляемое к знаменателю для повышения численной стабильности (по умолчанию: 1e-8)
- maximize (bool, необязательный) – максимизировать целевую функцию относительно параметров вместо ее минимизации (по умолчанию: False)
-
add_param_group(param_group)[исходный код] -
Добавить группу параметров в
Optimizersparam_groups.Это может быть полезно при дообучении предварительно обученной сети: замороженные слои можно сделать обучаемыми и добавлять в
Optimizerпо мере обучения.- Параметры:
-
param_group (dict) – задает, какие тензоры следует оптимизировать, а также параметры оптимизации, специфичные для группы.
-
load_state_dict(state_dict)[исходный код] -
Загрузить состояние оптимизатора.
- Параметры:
-
state_dict (dict) – состояние оптимизатора. Должно быть объектом, возвращенным вызовом
state_dict().
Предупреждение
Убедитесь, что этот метод вызывается после инициализации
torch.optim.lr_scheduler.LRScheduler, поскольку его вызов до инициализации перезапишет загруженные скорости обучения.Примечание
Имена параметров (если они присутствуют под ключом “param_names” в каждой группе параметров в
state_dict()) не влияют на процесс загрузки. Чтобы использовать имена параметров в особых случаях (например, когда параметры в загруженном словаре состояния отличаются от инициализированных в оптимизаторе), следует реализовать пользовательскийregister_load_state_dict_pre_hookдля соответствующей адаптации загружаемого словаря. Еслиparam_namesприсутствуют в загруженном словаре состоянияparam_groups, они будут сохранены и заменят текущие имена, если таковые имеются, в состоянии оптимизатора. Если их нет в загруженном словаре состояния,param_namesоптимизатора останутся без изменений.Пример
>>> optimizer = ... # initialized optimizer matching the saved state >>> scheduler1 = torch.optim.lr_scheduler.LinearLR( ... optimizer, ... start_factor=0.1, ... end_factor=1, ... total_iters=20, ... ) >>> scheduler2 = torch.optim.lr_scheduler.CosineAnnealingLR( ... optimizer, ... T_max=80, ... eta_min=3e-5, ... ) >>> lr = torch.optim.lr_scheduler.SequentialLR( ... optimizer, ... schedulers=[scheduler1, scheduler2], ... milestones=[20], ... ) >>> lr.load_state_dict(torch.load("./save_seq.pt")) >>> # now load the optimizer checkpoint after loading the LRScheduler >>> optimizer.load_state_dict(torch.load("./save_optim.pt"))
-
register_load_state_dict_post_hook(hook, prepend=False)[исходный код] -
Зарегистрировать постперехватчик load_state_dict, который будет вызван после вызова
load_state_dict(). Он должен иметь следующую сигнатуру:hook(optimizer) -> None
Аргумент
optimizer— это используемый экземпляр оптимизатора.После вызова
load_state_dictдляselfперехватчик будет вызван с аргументомself. Зарегистрированный перехватчик можно использовать для постобработки после того, какload_state_dictзагрузитstate_dict.- Параметры:
-
- hook (Callable) – определенный пользователем перехватчик, который нужно зарегистрировать.
-
prepend (bool) – Если значение True, предоставленный постперехватчик
hookбудет вызван до всех уже зарегистрированных постперехватчиков дляload_state_dict. В противном случае предоставленныйhookбудет вызван после всех уже зарегистрированных постперехватчиков. (по умолчанию: False)
- Возвращает:
-
объект, с помощью которого можно удалить добавленный перехватчик, вызвав
handle.remove() - Тип возвращаемого значения:
-
torch.utils.hooks.RemovableHandle
-
register_load_state_dict_pre_hook(hook, prepend=False)[исходный код] -
Зарегистрировать предперехватчик load_state_dict, который будет вызван перед вызовом
load_state_dict(). Он должен иметь следующую сигнатуру:hook(optimizer, state_dict) -> state_dict or None
Аргумент
optimizer— это используемый экземпляр оптимизатора, а аргументstate_dict— это поверхностная копияstate_dict, переданного пользователем вload_state_dict. Перехватчик может изменить state_dict на месте или, при желании, вернуть новый. Если возвращается state_dict, он будет использован для загрузки в оптимизатор.Перед вызовом
load_state_dictдляselfперехватчик будет вызван с аргументамиselfиstate_dict. Зарегистрированный перехватчик можно использовать для предварительной обработки перед вызовомload_state_dict.- Параметры:
-
- hook (Callable) – определенный пользователем перехватчик, который нужно зарегистрировать.
-
prepend (bool) – Если значение True, предоставленный предперехватчик
hookбудет вызван до всех уже зарегистрированных предперехватчиков дляload_state_dict. В противном случае предоставленныйhookбудет вызван после всех уже зарегистрированных предперехватчиков. (по умолчанию: False)
- Возвращает:
-
объект, с помощью которого можно удалить добавленный перехватчик, вызвав
handle.remove() - Тип возвращаемого значения:
-
torch.utils.hooks.RemovableHandle
-
register_state_dict_post_hook(hook, prepend=False)[исходный код] -
Зарегистрировать постперехватчик словаря состояния, который будет вызван после вызова
state_dict().Он должен иметь следующую сигнатуру:
hook(optimizer, state_dict) -> state_dict or None
После создания
state_dictдляselfперехватчик будет вызван с аргументамиselfиstate_dict. Перехватчик может изменить state_dict на месте или, при желании, вернуть новый. Зарегистрированный перехватчик можно использовать для постобработкиstate_dictперед его возвратом.- Параметры:
-
- hook (Callable) – определенный пользователем перехватчик, который нужно зарегистрировать.
-
prepend (bool) – Если значение True, предоставленный постперехватчик
hookбудет вызван до всех уже зарегистрированных постперехватчиков дляstate_dict. В противном случае предоставленныйhookбудет вызван после всех уже зарегистрированных постперехватчиков. (по умолчанию: False)
- Возвращает:
-
объект, с помощью которого можно удалить добавленный перехватчик, вызвав
handle.remove() - Тип возвращаемого значения:
-
torch.utils.hooks.RemovableHandle
-
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.RemovableHandle
-
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 -
группа параметров представляет собой словарь. Каждая группа параметров содержит метаданные, специфичные для оптимизатора, такие как скорость обучения и затухание весов, а также список идентификаторов параметров, входящих в группу. Если группа параметров была инициализирована с помощью
named_parameters(), имена также будут сохранены в словаре состояния.
-
ПРИМЕЧАНИЕ: Идентификаторы параметров могут выглядеть как индексы, но это всего лишь идентификаторы, связывающие состояние с param_group. При загрузке из state_dict оптимизатор сопоставляет элементы param_group
params(целочисленные идентификаторы) иparam_groupsоптимизатора (фактическиеnn.Parameters), чтобы сопоставить состояния БЕЗ дополнительной проверки.Возвращаемый словарь состояния может выглядеть примерно так:
{ '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] 'param_names' ['param0'] (optional) }, { 'lr': 0.001, 'weight_decay': 0.5, ... 'params': [1, 2, 3] 'param_names': ['param1', 'layer.weight', 'layer.bias'] (optional) } ] } -
-
step(closure=None)[исходный код] -
Выполнить один шаг оптимизации.
- Параметры:
-
closure (Callable, необязательный) – замыкание, повторно вычисляющее модель и возвращающее значение функции потерь.
-
zero_grad(set_to_none=True)[исходный код] -
Сбросить градиенты всех оптимизируемых тензоров
torch.Tensor.- Параметры:
-
set_to_none (bool, необязательный) –
Вместо обнуления установить градиенты в None. По умолчанию:
TrueКак правило, это уменьшает объем используемой памяти и может немного повысить производительность. Однако это меняет некоторые варианты поведения. Например:
- Когда пользователь пытается получить доступ к градиенту и выполнить над ним операции вручную, атрибут None и тензор, заполненный нулями, ведут себя по-разному.
- Если пользователь запрашивает
zero_grad(set_to_none=True), а затем выполняет обратный проход, для параметров, не получивших градиент, гарантируется значение None у.grads. - Оптимизаторы
torch.optimведут себя по-разному, когда градиент равен 0 или None (в одном случае выполняется шаг с градиентом 0, а в другом шаг полностью пропускается).
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.optim.SparseAdam.html