AveragedModel
-
class torch.optim.swa_utils.AveragedModel(model, device=None, avg_fn=None, multi_avg_fn=None, use_buffers=False)[исходный код] -
Реализует усреднение модели для стохастического усреднения весов (SWA) и экспоненциального скользящего среднего (EMA).
Метод стохастического усреднения весов предложен в работе Усреднение весов приводит к более широким минимумам и лучшей обобщающей способности Павла Измайлова, Дмитрия Подоприхина, Тимура Гарипова, Дмитрия Ветрова и Эндрю Гордона Уилсона (UAI 2018).
Экспоненциальное скользящее среднее — это вариант усреднения Поляка, в котором вместо равных весов для итераций используются экспоненциальные веса.
Класс AveragedModel создает копию предоставленного модуля
modelна устройствеdeviceи позволяет вычислять скользящие средние параметровmodel.- Параметры:
-
- model (torch.nn.Module) – модель для использования с SWA/EMA
-
device (torch.device, необязательно) – если указано, усредненная модель будет храниться на
device -
avg_fn (функция, необязательно) – функция усреднения, используемая для обновления параметров; функция должна принимать текущее значение параметра
AveragedModel, текущее значение параметраmodelи количество уже усредненных моделей; если значение равно None, используется среднее с равными весами (по умолчанию: None) -
multi_avg_fn (функция, необязательно) – функция усреднения, используемая для обновления параметров на месте; функция должна принимать текущие значения параметров
AveragedModelв виде списка, текущие значения параметровmodelв виде списка и количество уже усредненных моделей; если значение равно None, используется среднее с равными весами (по умолчанию: None) -
use_buffers (bool) – если
True, скользящие средние будут вычисляться как для параметров, так и для буферов модели. (по умолчанию:False)
Пример
>>> 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 i in range(300): >>> for input, target in loader: >>> optimizer.zero_grad() >>> loss_fn(model(input), target).backward() >>> optimizer.step() >>> if i > 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)
Вы также можете использовать собственные функции усреднения с параметрами
avg_fnилиmulti_avg_fn. Если функция усреднения не указана, по умолчанию вычисляется среднее весов с равными коэффициентами (SWA).Пример
>>> # Compute exponential moving averages of the weights and buffers >>> ema_model = torch.optim.swa_utils.AveragedModel(model, >>> torch.optim.swa_utils.get_ema_multi_avg_fn(0.9), use_buffers=True)
Примечание
При использовании SWA/EMA с моделями, содержащими пакетную нормализацию, может потребоваться обновить статистики активаций для пакетной нормализации. Это можно сделать с помощью
torch.optim.swa_utils.update_bn()или установивuse_buffersвTrue. Первый способ обновляет статистики на этапе после обучения, передавая данные через модель. Второй выполняет это во время обновления параметров, усредняя все буферы. Эмпирические данные показывают, что обновление статистик в слоях нормализации повышает точность, однако рекомендуется экспериментально проверить, какой способ дает наилучшие результаты в вашей задаче.Примечание
avg_fnиmulti_avg_fnне сохраняются вstate_dict()модели.Примечание
При первом вызове
update_parameters()(то есть когдаn_averagedравно0) параметрыmodelкопируются в параметрыAveragedModel. При каждом последующем вызовеupdate_parameters()для обновления параметров используется функцияavg_fn.-
add_module(name, module)[исходный код] -
Добавляет дочерний модуль к текущему модулю.
К модулю можно обратиться как к атрибуту, используя указанное имя.
-
apply(fn)[исходный код] -
Рекурсивно применяет
fnк каждому подмодулю (возвращаемому.children()), а также к самому модулю.Обычно этот метод используется для инициализации параметров модели (см. также torch.nn.init).
- Параметры:
-
fn (
Module-> None) – функция, применяемая к каждому подмодулю - Возвращает:
-
self
- Тип возвращаемого значения:
Пример:
>>> @torch.no_grad() >>> def init_weights(m): >>> print(m) >>> if type(m) is nn.Linear: >>> m.weight.fill_(1.0) >>> print(m.weight) >>> net = nn.Sequential(nn.Linear(2, 2), nn.Linear(2, 2)) >>> net.apply(init_weights) Linear(in_features=2, out_features=2, bias=True) Parameter containing: tensor([[1., 1.], [1., 1.]], requires_grad=True) Linear(in_features=2, out_features=2, bias=True) Parameter containing: tensor([[1., 1.], [1., 1.]], requires_grad=True) Sequential( (0): Linear(in_features=2, out_features=2, bias=True) (1): Linear(in_features=2, out_features=2, bias=True) )
-
bfloat16()[исходный код] -
Преобразует все параметры с плавающей точкой и буферы к типу данных
bfloat16.Примечание
Этот метод изменяет модуль на месте.
- Возвращает:
-
self
- Тип возвращаемого значения:
-
buffers(recurse=True)[исходный код] -
Возвращает итератор по буферам модуля.
- Параметры:
-
recurse (bool) – если True, возвращает буферы этого модуля и всех его подмодулей. В противном случае возвращает только буферы, являющиеся непосредственными членами этого модуля.
- Возвращает значения:
-
torch.Tensor – буфер модуля
- Тип возвращаемого значения:
Пример:
>>> for buf in model.buffers(): >>> print(type(buf), buf.size()) <class 'torch.Tensor'> (20L,) <class 'torch.Tensor'> (20L, 1L, 5L, 5L)
-
children()[исходный код] -
Возвращает итератор по непосредственным дочерним модулям.
-
compile(*args, **kwargs)[исходный код] -
Компилирует метод forward этого модуля с помощью
torch.compile().Метод
__call__этого модуля компилируется, а все аргументы без изменений передаются вtorch.compile().Подробные сведения об аргументах этой функции см. в разделе
torch.compile().
-
cpu()[исходный код] -
Перемещает все параметры и буферы модели на CPU.
Примечание
Этот метод изменяет модуль на месте.
- Возвращает:
-
self
- Тип возвращаемого значения:
-
cuda(device=None)[исходный код] -
Перемещает все параметры и буферы модели на GPU.
При этом связанные параметры и буферы также становятся другими объектами. Поэтому этот метод следует вызывать до создания оптимизатора, если модуль будет находиться на GPU во время оптимизации.
Примечание
Этот метод изменяет модуль на месте.
-
double()[исходный код] -
Преобразует все параметры с плавающей точкой и буферы к типу данных
double.Примечание
Этот метод изменяет модуль на месте.
- Возвращает:
-
self
- Тип возвращаемого значения:
-
eval()[исходный код] -
Переводит модуль в режим оценки.
Это влияет только на некоторые модули. Подробные сведения о поведении конкретных модулей в режимах обучения и оценки, то есть о том, затрагивает ли их этот метод, см. в документации соответствующих модулей; например,
Dropout,BatchNormи т. д.Это эквивалентно вызову
self.train(False).Сравнение
.eval()с несколькими похожими механизмами, которые можно с ним перепутать, см. в разделе Локальное отключение вычисления градиентов.- Возвращает:
-
self
- Тип возвращаемого значения:
-
extra_repr()[исходный код] -
Возвращает дополнительное представление модуля.
Чтобы выводить дополнительную информацию в собственном формате, переопределите этот метод в своих модулях. Допускаются как однострочные, так и многострочные строки.
- Тип возвращаемого значения:
-
float()[исходный код] -
Преобразует все параметры с плавающей точкой и буферы к типу данных
float.Примечание
Этот метод изменяет модуль на месте.
- Возвращает:
-
self
- Тип возвращаемого значения:
-
forward(*args, **kwargs)[исходный код] -
Прямой проход.
-
get_buffer(target)[исходный код] -
Возвращает буфер, указанный в
target, если он существует; в противном случае вызывает ошибку.Более подробное описание функциональности этого метода и правильного указания
targetсм. в документации кget_submodule.- Параметры:
-
target (str) – полное строковое имя буфера, который необходимо найти. (См.
get_submodule, где описано, как указать полную строку.) - Возвращает:
-
Буфер, указанный в
target - Тип возвращаемого значения:
- Вызывает исключение:
-
AttributeError – если целевая строка указывает на недопустимый путь или разрешается в объект, который не является буфером
-
get_extra_state()[исходный код] -
Возвращает любые дополнительные данные состояния, которые нужно включить в state_dict модуля.
Если вашему модулю необходимо хранить дополнительные данные состояния, реализуйте этот метод и соответствующий ему метод
set_extra_state(). Эта функция вызывается при созданииstate_dict()модуля.Обратите внимание: дополнительные данные состояния должны поддерживать сериализацию с помощью pickle, чтобы сериализация state_dict работала корректно. Гарантии обратной совместимости предоставляются только для сериализации тензоров; изменение сериализованной формы pickle других объектов может нарушить обратную совместимость.
- Возвращает:
-
Любые дополнительные данные состояния, сохраняемые в state_dict модуля
- Тип возвращаемого значения:
-
get_parameter(target)[исходный код] -
Возвращает параметр, указанный в
target, если он существует; в противном случае вызывает ошибку.Более подробное описание функциональности этого метода и правильного указания
targetсм. в документации кget_submodule.- Параметры:
-
target (str) – полное строковое имя параметра, который необходимо найти. (См.
get_submodule, где описано, как указать полную строку.) - Возвращает:
-
Параметр, указанный в
target - Тип возвращаемого значения:
-
torch.nn.Parameter
- Вызывает исключение:
-
AttributeError – если целевая строка указывает на недопустимый путь или разрешается в объект, который не является
nn.Parameter
-
get_submodule(target)[исходный код] -
Возвращает подмодуль, указанный в
target, если он существует; в противном случае вызывает ошибку.Например, предположим, что у вас есть
nn.ModuleAследующей структуры:A( (net_b): Module( (net_c): Module( (conv): Conv2d(16, 33, kernel_size=(3, 3), stride=(2, 2)) ) (linear): Linear(in_features=100, out_features=200, bias=True) ) )(На схеме показан
nn.ModuleA.Aимеет вложенный подмодульnet_b, у которого, в свою очередь, есть два подмодуля:net_cиlinear. Затем уnet_cесть подмодульconv.)Чтобы проверить, есть ли у нас подмодуль
linear, нужно вызватьget_submodule("net_b.linear"). Чтобы проверить наличие подмодуляconv, нужно вызватьget_submodule("net_b.net_c.conv").Время выполнения
get_submoduleограничено глубиной вложенности модулей вtarget. Запрос кnamed_modulesдает тот же результат, но его сложность составляет O(N) по числу транзитивных модулей. Поэтому для простой проверки наличия подмодуля всегда следует использоватьget_submodule.- Параметры:
-
target (str) – полное строковое имя искомого подмодуля. (О том, как указать полную строку, см. пример выше.)
- Возвращает:
-
Подмодуль, указанный в
target - Тип возвращаемого значения:
- Вызывает исключение:
-
AttributeError – если на любом этапе пути, заданного целевой строкой, под путь разрешается в несуществующее имя атрибута или в объект, который не является экземпляром
nn.Module.
-
half()[исходный код] -
Преобразует все параметры с плавающей точкой и буферы к типу данных
half.Примечание
Этот метод изменяет модуль на месте.
- Возвращает:
-
self
- Тип возвращаемого значения:
-
ipu(device=None)[исходный код] -
Перемещает все параметры и буферы модели на IPU.
При этом связанные параметры и буферы также становятся другими объектами. Поэтому этот метод следует вызывать до создания оптимизатора, если модуль будет находиться на IPU во время оптимизации.
Примечание
Этот метод изменяет модуль на месте.
-
load_state_dict(state_dict, strict=True, assign=False)[исходный код] -
Скопировать параметры и буферы из
state_dictв этот модуль и его дочерние модули.Если
strictравноTrue, ключиstate_dictдолжны в точности совпадать с ключами, возвращаемыми функциейstate_dict()этого модуля.Предупреждение
Если
assignравноTrue, оптимизатор необходимо создать после вызоваload_state_dict, если толькоget_swap_module_params_on_conversion()не равноTrue.- Параметры:
-
- state_dict (dict) – словарь, содержащий параметры и постоянные буферы.
-
strict (bool, optional) – строго ли проверять, что ключи в
state_dictсовпадают с ключами, возвращаемыми функциейstate_dict()этого модуля. По умолчанию:True -
assign (bool, optional) – если задано значение
False, свойства тензоров текущего модуля сохраняются; если задано значениеTrue, сохраняются свойства тензоров в словаре состояния. Единственное исключение — полеrequires_gradобъектаParameter, для которого сохраняется значение из модуля. По умолчанию:False
- Возвращает:
-
-
-
missing_keys is a list of str containing any keys that are expected -
в этом модуле, но отсутствующие в предоставленном
state_dict.
-
-
-
unexpected_keys is a list of str containing the keys that are not -
ожидаемые этим модулем, но присутствующие в предоставленном
state_dict.
-
-
- Тип возвращаемого значения:
-
NamedTupleс полямиmissing_keysиunexpected_keys
Примечание
Если параметр или буфер зарегистрирован как
Noneи соответствующий ему ключ существует вstate_dict, вызовload_state_dict()вызовет исключениеRuntimeError.
-
modules(remove_duplicate=True)[исходный код] -
Вернуть итератор по всем модулям сети.
- Параметры:
-
remove_duplicate (bool) – удалять ли повторяющиеся экземпляры модулей из результата.
- Выдаёт:
-
Module – модуль сети
- Тип возвращаемого значения:
Примечание
По умолчанию повторяющиеся модули возвращаются только один раз. В следующем примере
lбудет возвращён только один раз.Пример:
>>> l = nn.Linear(2, 2) >>> net = nn.Sequential(l, l) >>> for idx, m in enumerate(net.modules()): ... print(idx, '->', m) 0 -> Sequential( (0): Linear(in_features=2, out_features=2, bias=True) (1): Linear(in_features=2, out_features=2, bias=True) ) 1 -> Linear(in_features=2, out_features=2, bias=True)
-
mtia(device=None)[исходный код] -
Переместить все параметры и буферы модели на MTIA.
При этом связанные параметры и буферы также становятся другими объектами. Поэтому этот метод следует вызывать до создания оптимизатора, если модуль будет находиться на MTIA во время оптимизации.
Примечание
Этот метод изменяет модуль на месте.
-
named_buffers(prefix='', recurse=True, remove_duplicate=True)[исходный код] -
Вернуть итератор по буферам модуля, выдавая как имя буфера, так и сам буфер.
- Параметры:
-
- prefix (str) – префикс, добавляемый ко всем именам буферов.
- recurse (bool, optional) – если True, выдаёт буферы этого модуля и всех его подмодулей. В противном случае выдаёт только буферы, непосредственно принадлежащие этому модулю. По умолчанию True.
- remove_duplicate (bool, optional) – удалять ли повторяющиеся буферы из результата. По умолчанию True.
- Выдаёт:
-
(str, torch.Tensor) – кортеж, содержащий имя и буфер
- Тип возвращаемого значения:
Пример:
>>> for name, buf in self.named_buffers(): >>> if name in ['running_var']: >>> print(buf.size())
-
named_children()[исходный код] -
Вернуть итератор по непосредственным дочерним модулям, выдавая как имя модуля, так и сам модуль.
- Выдаёт:
-
(str, Module) – кортеж, содержащий имя и дочерний модуль
- Тип возвращаемого значения:
Пример:
>>> for name, module in model.named_children(): >>> if name in ['conv4', 'conv5']: >>> print(module)
-
named_modules(memo=None, prefix='', remove_duplicate=True)[исходный код] -
Вернуть итератор по всем модулям сети, выдавая как имя модуля, так и сам модуль.
- Параметры:
- Выдаёт:
-
(str, Module) – кортеж из имени и модуля
Примечание
Повторяющиеся модули возвращаются только один раз. В следующем примере
lбудет возвращён только один раз.Пример:
>>> l = nn.Linear(2, 2) >>> net = nn.Sequential(l, l) >>> for idx, m in enumerate(net.named_modules()): ... print(idx, '->', m) 0 -> ('', Sequential( (0): Linear(in_features=2, out_features=2, bias=True) (1): Linear(in_features=2, out_features=2, bias=True) )) 1 -> ('0', Linear(in_features=2, out_features=2, bias=True))
-
named_parameters(prefix='', recurse=True, remove_duplicate=True)[исходный код] -
Вернуть итератор по параметрам модуля, выдавая как имя параметра, так и сам параметр.
- Параметры:
-
- prefix (str) – префикс, добавляемый ко всем именам параметров.
- recurse (bool) – если True, выдаёт параметры этого модуля и всех его подмодулей. В противном случае выдаёт только параметры, непосредственно принадлежащие этому модулю.
- remove_duplicate (bool, optional) – удалять ли повторяющиеся параметры из результата. По умолчанию True.
- Выдаёт:
-
(str, Parameter) – кортеж, содержащий имя и параметр
- Тип возвращаемого значения:
Пример:
>>> for name, param in self.named_parameters(): >>> if name in ['bias']: >>> print(param.size())
-
parameters(recurse=True)[исходный код] -
Вернуть итератор по параметрам модуля.
Точный порядок возвращаемых параметров не определён, однако повторные вызовы метода
parameters()для неизменённого модуля возвращают параметры в одном и том же порядке.Обычно этот итератор передают оптимизатору.
- Параметры:
-
recurse (bool) – если True, выдаёт параметры этого модуля и всех его подмодулей. В противном случае выдаёт только параметры, непосредственно принадлежащие этому модулю.
- Выдаёт:
-
Parameter – параметр модуля
- Тип возвращаемого значения:
Пример:
>>> for param in model.parameters(): >>> print(type(param), param.size()) <class 'torch.Tensor'> (20L,) <class 'torch.Tensor'> (20L, 1L, 5L, 5L)
-
register_backward_hook(hook)[исходный код] -
Зарегистрировать хук обратного прохода для модуля.
Эта функция устарела; вместо неё рекомендуется использовать
register_full_backward_hook(). Поведение этой функции изменится в будущих версиях.- Возвращает:
-
обработчик, с помощью которого можно удалить добавленный хук, вызвав
handle.remove() - Тип возвращаемого значения:
-
torch.utils.hooks.RemovableHandle
-
register_buffer(name, tensor, persistent=True)[исходный код] -
Добавить буфер в модуль.
Обычно этот метод используется для регистрации буфера, который не следует считать параметром модели. Например,
running_meanBatchNorm не является параметром, но входит в состояние модуля. По умолчанию буферы являются постоянными и сохраняются вместе с параметрами. Это поведение можно изменить, задав дляpersistentзначениеFalse. Единственное различие между постоянным и непостоянным буфером состоит в том, что последний не входит вstate_dictэтого модуля.К буферам можно обращаться как к атрибутам, используя заданные имена.
- Параметры:
-
- name (str) – имя буфера. К буферу можно обратиться из этого модуля по заданному имени
-
tensor (Tensor or None) – регистрируемый буфер. Если
None, операции над буферами, такие какcuda, игнорируются. ЕслиNone, буфер не включается вstate_dictмодуля. -
persistent (bool) – входит ли буфер в
state_dictэтого модуля.
Пример:
>>> self.register_buffer('running_mean', torch.zeros(num_features))
-
register_forward_hook(hook, *, prepend=False, with_kwargs=False, always_call=False)[исходный код] -
Зарегистрировать хук прямого прохода для модуля.
Хук будет вызываться каждый раз после того, как
forward()вычислит результат.Если
with_kwargsравноFalseили не указано, входные данные содержат только позиционные аргументы, переданные модулю. Именованные аргументы не будут передаваться хукам, а будут передаваться толькоforward. Хук может изменять выходные данные. Он может изменять входные данные на месте, но это не повлияет на прямой проход, поскольку хук вызывается после вызоваforward(). Хук должен иметь следующую сигнатуру:hook(module, args, output) -> None or modified output
Если
with_kwargsравноTrue, хук прямого прохода получитkwargs, переданные функции прямого прохода, и должен вернуть выходные данные, возможно, изменённые. Хук должен иметь следующую сигнатуру:hook(module, args, kwargs, output) -> None or modified output
- Параметры:
-
- hook (Callable) – пользовательский хук, который необходимо зарегистрировать.
-
prepend (bool) – если
True, предоставленныйhookбудет вызван перед всеми существующими хукамиforwardэтогоtorch.nn.Module. В противном случае предоставленныйhookбудет вызван после всех существующих хуковforwardэтогоtorch.nn.Module. Обратите внимание, что глобальные хукиforward, зарегистрированные с помощьюregister_module_forward_hook(), будут вызваны перед всеми хуками, зарегистрированными этим методом. По умолчанию:False -
with_kwargs (bool) – если
True, хукhookполучит именованные аргументы, переданные функции прямого прохода. По умолчанию:False -
always_call (bool) – если
True, хукhookбудет выполнен независимо от того, возникнет ли исключение при вызове Module. По умолчанию:False
- Возвращает:
-
обработчик, с помощью которого можно удалить добавленный хук, вызвав
handle.remove() - Тип возвращаемого значения:
-
torch.utils.hooks.RemovableHandle
-
register_forward_pre_hook(hook, *, prepend=False, with_kwargs=False)[исходный код] -
Зарегистрировать предварительный хук прямого прохода для модуля.
Хук будет вызываться каждый раз перед вызовом
forward().Если
with_kwargsравно false или не указано, входные данные содержат только позиционные аргументы, переданные модулю. Именованные аргументы не будут передаваться хукам, а будут передаваться толькоforward. Хук может изменять входные данные. Из хука можно вернуть либо кортеж, либо одно изменённое значение. Если возвращается одиночное значение (если оно ещё не является кортежем), оно будет обёрнуто в кортеж. Хук должен иметь следующую сигнатуру:hook(module, args) -> None or modified input
Если
with_kwargsравно true, предварительный хук прямого прохода получит именованные аргументы, переданные функции прямого прохода. Если хук изменяет входные данные, необходимо вернуть и args, и kwargs. Хук должен иметь следующую сигнатуру:hook(module, args, kwargs) -> None or a tuple of modified input and kwargs
- Параметры:
-
- hook (Callable) – пользовательский хук, который необходимо зарегистрировать.
-
prepend (bool) – если true, предоставленный
hookбудет вызван перед всеми существующими хукамиforward_preэтогоtorch.nn.Module. В противном случае предоставленныйhookбудет вызван после всех существующих хуковforward_preэтогоtorch.nn.Module. Обратите внимание, что глобальные хукиforward_pre, зарегистрированные с помощьюregister_module_forward_pre_hook(), будут вызваны перед всеми хуками, зарегистрированными этим методом. По умолчанию:False -
with_kwargs (bool) – если true, хук
hookполучит именованные аргументы, переданные функции прямого прохода. По умолчанию:False
- Возвращает:
-
обработчик, с помощью которого можно удалить добавленный хук, вызвав
handle.remove() - Тип возвращаемого значения:
-
torch.utils.hooks.RemovableHandle
-
register_full_backward_hook(hook, prepend=False)[исходный код] -
Зарегистрировать хук обратного прохода для модуля.
Хук будет вызываться каждый раз при вычислении градиентов относительно модуля; правила его срабатывания следующие:
- Обычно хук срабатывает при вычислении градиентов относительно входных данных модуля.
- Если ни для одного из входных данных модуля не требуются градиенты, хук сработает при вычислении градиентов относительно выходных данных модуля.
- Если ни для одного из выходных данных модуля не требуются градиенты, хуки не сработают.
Хук должен иметь следующую сигнатуру:
hook(module, grad_input, grad_output) -> tuple(Tensor) or None
grad_inputиgrad_output— это кортежи, содержащие градиенты относительно входных и выходных данных соответственно. Хук не должен изменять свои аргументы, но при необходимости может возвращать новый градиент относительно входных данных, который будет использоваться вместоgrad_inputв последующих вычислениях.grad_inputсоответствует только входным данным, переданным в виде позиционных аргументов; все именованные аргументы игнорируются. Элементыgrad_inputиgrad_outputбудут иметь значениеNoneдля всех аргументов, не являющихся тензорами.По техническим причинам при применении этого хука к Module его функция прямого прохода получает представление каждого тензора, переданного модулю. Аналогично, вызывающая сторона получает представление каждого тензора, возвращённого функцией прямого прохода модуля.
Предупреждение
При использовании хуков обратного прохода изменение входных или выходных данных на месте не допускается и приведёт к ошибке.
- Параметры:
-
- hook (Callable) – пользовательский хук, который необходимо зарегистрировать.
-
prepend (bool) – если true, предоставленный
hookбудет вызван перед всеми существующими хукамиbackwardэтогоtorch.nn.Module. В противном случае предоставленныйhookбудет вызван после всех существующих хуковbackwardэтогоtorch.nn.Module. Обратите внимание, что глобальные хукиbackward, зарегистрированные с помощьюregister_module_full_backward_hook(), будут вызваны перед всеми хуками, зарегистрированными этим методом.
- Возвращает:
-
обработчик, с помощью которого можно удалить добавленный хук, вызвав
handle.remove() - Тип возвращаемого значения:
-
torch.utils.hooks.RemovableHandle
-
-
register_full_backward_pre_hook(hook, prepend=False)[исходный код] -
Зарегистрировать предварительный хук обратного прохода для модуля.
Хук будет вызываться каждый раз при вычислении градиентов модуля. Сигнатура хука должна быть следующей:
hook(module, grad_output) -> tuple[Tensor, ...], Tensor or None
Аргумент
grad_outputпредставляет собой кортеж. Хук не должен изменять свои аргументы, но может при необходимости вернуть новый градиент относительно выходных данных, который будет использоваться вместоgrad_outputв последующих вычислениях. Элементы вgrad_outputбудут равныNoneдля всех аргументов, не являющихся тензорами.По техническим причинам при применении этого хука к Module его функция forward получает представление каждого тензора, переданного модулю. Аналогично вызывающий код получает представление каждого тензора, возвращённого функцией forward модуля.
Предупреждение
Изменять входные данные на месте при использовании хуков обратного прохода нельзя; это приведёт к ошибке.
- Параметры:
-
- hook (Callable) – Пользовательский хук, который необходимо зарегистрировать.
-
prepend (bool) – Если значение истинно, предоставленный
hookбудет вызван до всех существующих хуковbackward_preэтогоtorch.nn.Module. В противном случае предоставленныйhookбудет вызван после всех существующих хуковbackward_preэтогоtorch.nn.Module. Обратите внимание, что глобальные хукиbackward_pre, зарегистрированные с помощьюregister_module_full_backward_pre_hook(), будут вызываться до всех хуков, зарегистрированных этим методом.
- Возвращает:
-
дескриптор, который можно использовать для удаления добавленного хука, вызвав
handle.remove() - Тип возвращаемого значения:
-
torch.utils.hooks.RemovableHandle
-
register_load_state_dict_post_hook(hook)[исходный код] -
Зарегистрировать пост-хук, который будет выполняться после вызова
load_state_dict()модуля.- Он должен иметь следующую сигнатуру::
-
hook(module, incompatible_keys) -> None
Аргумент
module— это текущий модуль, для которого зарегистрирован этот хук, а аргументincompatible_keys— этоNamedTuple, состоящий из атрибутовmissing_keysиunexpected_keys.missing_keys— этоlistизstr, содержащий отсутствующие ключи, аunexpected_keys— этоlistизstr, содержащий неожиданные ключи.При необходимости указанный incompatible_keys можно изменить на месте.
Обратите внимание, что проверки, выполняемые при вызове
load_state_dict()сstrict=True, зависят от изменений, внесённых хуком вmissing_keysилиunexpected_keys, как и ожидается. Добавление ключей в любой из наборов приведёт к возникновению ошибки приstrict=True, а удаление всех отсутствующих и неожиданных ключей позволит избежать ошибки.- Возвращает:
-
дескриптор, который можно использовать для удаления добавленного хука, вызвав
handle.remove() - Тип возвращаемого значения:
-
torch.utils.hooks.RemovableHandle
-
register_load_state_dict_pre_hook(hook)[исходный код] -
Зарегистрировать предварительный хук, который будет выполняться перед вызовом
load_state_dict()модуля.- Он должен иметь следующую сигнатуру::
-
hook(module, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs) -> None # noqa: B950
- Параметры:
-
hook (Callable) – Вызываемый хук, который будет вызван перед загрузкой словаря состояния.
-
register_module(name, module)[исходный код] -
Псевдоним для
add_module().
-
register_parameter(name, param)[исходный код] -
Добавить параметр в модуль.
К параметру можно обращаться как к атрибуту, используя указанное имя.
- Параметры:
-
- name (str) – имя параметра. К параметру можно обращаться из этого модуля, используя указанное имя
-
param (Parameter or None) – параметр, добавляемый в модуль. Если
None, операции над параметрами, такие какcuda, игнорируются. ЕслиNone, параметр не включается вstate_dictмодуля.
-
register_state_dict_post_hook(hook)[исходный код] -
Зарегистрировать пост-хук для метода
state_dict().- Он должен иметь следующую сигнатуру::
-
hook(module, state_dict, prefix, local_metadata) -> None
Зарегистрированные хуки могут изменять
state_dictна месте.
-
register_state_dict_pre_hook(hook)[исходный код] -
Зарегистрировать предварительный хук для метода
state_dict().- Он должен иметь следующую сигнатуру::
-
hook(module, prefix, keep_vars) -> None
Зарегистрированные хуки можно использовать для предварительной обработки перед вызовом
state_dict.
-
requires_grad_(requires_grad=True)[исходный код] -
Изменить, должна ли autograd записывать операции с параметрами этого модуля.
Этот метод изменяет атрибуты
requires_gradпараметров на месте.Этот метод полезен для заморозки части модуля при дообучении или для раздельного обучения частей модели (например, при обучении GAN).
См. раздел Локальное отключение вычисления градиентов, где сравнивается
.requires_grad_()с несколькими похожими механизмами, которые можно с ним спутать.
-
set_extra_state(state)[исходный код] -
Задать дополнительные данные состояния, содержащиеся в загруженном
state_dict.Эта функция вызывается из
load_state_dict()для обработки дополнительных данных состояния, найденных вstate_dict. Реализуйте эту функцию и соответствующую ейget_extra_state()для своего модуля, если необходимо хранить дополнительные данные состояния в егоstate_dict.- Параметры:
-
state (dict) – Дополнительные данные состояния из
state_dict
-
set_submodule(target, module, strict=False)[исходный код] -
Задать подмодуль, указанный в
target, если он существует; в противном случае вызвать ошибку.Примечание
Если для
strictзадано значениеFalse(по умолчанию), метод заменит существующий подмодуль или создаст новый, если родительский модуль существует. Если дляstrictзадано значениеTrue, метод попытается только заменить существующий подмодуль и вызовет ошибку, если подмодуля не существует.Например, предположим, что у вас есть
nn.ModuleAследующего вида:A( (net_b): Module( (net_c): Module( (conv): Conv2d(3, 3, 3) ) (linear): Linear(3, 3) ) )(На диаграмме показан
nn.ModuleA.Aсодержит вложенный подмодульnet_b, в котором, в свою очередь, есть два подмодуля:net_cиlinear. Затемnet_cсодержит подмодульconv.)Чтобы заменить
Conv2dновым подмодулемLinear, можно вызватьset_submodule("net_b.net_c.conv", nn.Linear(1, 1)), гдеstrictможет бытьTrueилиFalseЧтобы добавить новый подмодуль
Conv2dв существующий модульnet_b, вызовитеset_submodule("net_b.conv", nn.Conv2d(1, 1, 1)).Если в приведённом выше примере задать
strict=Trueи вызватьset_submodule("net_b.conv", nn.Conv2d(1, 1, 1), strict=True), будет вызвано исключение AttributeError, поскольку вnet_bнет подмодуля с именемconv.- Параметры:
-
- target (str) – Полное строковое имя подмодуля, который нужно найти. (См. пример выше, где показано, как указать полное строковое имя.)
- module (Module) – Модуль, который необходимо задать в качестве подмодуля.
-
strict (bool) – Если
False, метод заменит существующий подмодуль или создаст новый, если родительский модуль существует. ЕслиTrue, метод попытается только заменить существующий подмодуль и вызовет ошибку, если подмодуль ещё не существует.
- Вызывает исключения:
-
-
ValueError – Если строка
targetпуста илиmoduleне является экземпляромnn.Module. -
AttributeError – Если на любом этапе пути, заданного строкой
target, (под)путь разрешается в несуществующее имя атрибута или объект, не являющийся экземпляромnn.Module.
-
ValueError – Если строка
-
См.
torch.Tensor.share_memory_().- Тип возвращаемого значения:
-
Self
-
state_dict(*args, destination=None, prefix='', keep_vars=False)[исходный код] -
Вернуть словарь, содержащий ссылки на всё состояние модуля.
В него включаются как параметры, так и постоянные буферы (например, скользящие средние). Ключи соответствуют именам параметров и буферов. Параметры и буферы, заданные как
None, не включаются.Примечание
Возвращаемый объект является поверхностной копией. Он содержит ссылки на параметры и буферы модуля.
Предупреждение
В настоящее время
state_dict()также принимает позиционные аргументы дляdestination,prefixиkeep_varsв указанном порядке. Однако от этого способа использования отказываются, и в будущих выпусках будут поддерживаться только именованные аргументы.Предупреждение
Не используйте аргумент
destination, поскольку он не предназначен для конечных пользователей.- Параметры:
-
-
destination (dict, optional) – Если указан, состояние модуля будет записано в этот словарь, и будет возвращён тот же объект. В противном случае будет создан и возвращён
OrderedDict. Значение по умолчанию:None. -
prefix (str, optional) – префикс, добавляемый к именам параметров и буферов для формирования ключей в state_dict. Значение по умолчанию:
''. -
keep_vars (bool, optional) – по умолчанию объекты
Tensor, возвращаемые в словаре состояния, отсоединены от autograd. Если задано значениеTrue, отсоединение не выполняется. Значение по умолчанию:False.
-
destination (dict, optional) – Если указан, состояние модуля будет записано в этот словарь, и будет возвращён тот же объект. В противном случае будет создан и возвращён
- Возвращает:
-
словарь, содержащий полное состояние модуля
- Тип возвращаемого значения:
Пример:
>>> module.state_dict().keys() ['bias', 'weight']
-
to(*args, **kwargs)[исходный код] -
Переместить и/или привести параметры и буферы к нужному типу.
Метод можно вызвать следующим образом:
- to(device=None, dtype=None, non_blocking=False)[исходный код]
- to(dtype, non_blocking=False)[исходный код]
- to(tensor, non_blocking=False)[исходный код]
- to(memory_format=torch.channels_last)[исходный код]
Сигнатура метода аналогична сигнатуре
torch.Tensor.to(), но он принимает только числа с плавающей запятой или комплексныеdtype. Кроме того, этот метод приводит к типуdtype(если он задан) только параметры и буферы с плавающей запятой или комплексные. Целочисленные параметры и буферы будут перемещеныdevice, если он задан, но их типы данных останутся без изменений. Если заданоnon_blocking, метод по возможности пытается выполнить преобразование/перемещение асинхронно относительно хоста, например при перемещении тензоров CPU с закреплённой памятью на устройства CUDA.Примеры приведены ниже.
Примечание
Этот метод изменяет модуль на месте.
- Параметры:
-
-
device (
torch.device) – требуемое устройство для параметров и буферов этого модуля -
dtype (
torch.dtype) – требуемый тип данных с плавающей запятой или комплексный тип данных для параметров и буферов этого модуля - tensor (torch.Tensor) – тензор, тип данных и устройство которого будут использоваться как требуемые тип данных и устройство для всех параметров и буферов этого модуля
-
memory_format (
torch.memory_format) – требуемый формат памяти для четырёхмерных параметров и буферов этого модуля (только именованный аргумент)
-
device (
- Возвращает:
-
self
- Тип возвращаемого значения:
Примеры:
>>> linear = nn.Linear(2, 2) >>> linear.weight Parameter containing: tensor([[ 0.1913, -0.3420], [-0.5113, -0.2325]]) >>> linear.to(torch.double) Linear(in_features=2, out_features=2, bias=True) >>> linear.weight Parameter containing: tensor([[ 0.1913, -0.3420], [-0.5113, -0.2325]], dtype=torch.float64) >>> gpu1 = torch.device("cuda:1") >>> linear.to(gpu1, dtype=torch.half, non_blocking=True) Linear(in_features=2, out_features=2, bias=True) >>> linear.weight Parameter containing: tensor([[ 0.1914, -0.3420], [-0.5112, -0.2324]], dtype=torch.float16, device='cuda:1') >>> cpu = torch.device("cpu") >>> linear.to(cpu) Linear(in_features=2, out_features=2, bias=True) >>> linear.weight Parameter containing: tensor([[ 0.1914, -0.3420], [-0.5112, -0.2324]], dtype=torch.float16) >>> linear = nn.Linear(2, 2, bias=None).to(torch.cdouble) >>> linear.weight Parameter containing: tensor([[ 0.3741+0.j, 0.2382+0.j], [ 0.5593+0.j, -0.4443+0.j]], dtype=torch.complex128) >>> linear(torch.ones(3, 2, dtype=torch.cdouble)) tensor([[0.6122+0.j, 0.1150+0.j], [0.6122+0.j, 0.1150+0.j], [0.6122+0.j, 0.1150+0.j]], dtype=torch.complex128)
-
to_empty(*, device, recurse=True)[исходный код] -
Переместить параметры и буферы на указанное устройство без копирования хранилища.
- Параметры:
-
-
device (
torch.device) – требуемое устройство для параметров и буферов этого модуля. - recurse (bool) – нужно ли рекурсивно перемещать параметры и буферы подмодулей на указанное устройство.
-
device (
- Возвращает:
-
self
- Тип возвращаемого значения:
-
train(mode=True)[исходный код] -
Перевести модуль в режим обучения.
Это влияет только на некоторые модули. Подробные сведения о поведении отдельных модулей в режиме обучения/оценки см. в их документации, в частности о том, затрагиваются ли они, например,
Dropout,BatchNormи т. д.
-
type(dst_type)[исходный код] -
Привести все параметры и буферы к типу
dst_type.Примечание
Этот метод изменяет модуль на месте.
-
update_parameters(model)[исходный код] -
Обновить параметры модели.
-
xpu(device=None)[исходный код] -
Переместить все параметры и буферы модели на XPU.
При этом связанные параметры и буферы также становятся другими объектами. Поэтому, если модуль будет находиться на XPU во время оптимизации, этот метод следует вызывать до создания оптимизатора.
Примечание
Этот метод изменяет модуль на месте.
-
zero_grad(set_to_none=True)[исходный код] -
Сбросить градиенты всех параметров модели.
Для дополнительной информации см. аналогичную функцию в
torch.optim.Optimizer.- Параметры:
-
set_to_none (bool) – вместо установки нулевого значения установить для градиентов значение None. Подробности см. в
torch.optim.Optimizer.zero_grad().
-
© 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.swa_utils.AveragedModel.html