Spec-Zone.ru › PyTorch 2.14

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) [исходный код]

Добавляет дочерний модуль к текущему модулю.

К модулю можно обратиться как к атрибуту, используя указанное имя.

Параметры:
  • name (str) – имя дочернего модуля. К дочернему модулю можно обратиться из этого модуля, используя указанное имя
  • module (Module) – дочерний модуль, добавляемый к модулю.
apply(fn) [исходный код]

Рекурсивно применяет fn к каждому подмодулю (возвращаемому .children()), а также к самому модулю.

Обычно этот метод используется для инициализации параметров модели (см. также torch.nn.init).

Параметры:

fn (Module -> None) – функция, применяемая к каждому подмодулю

Возвращает:

self

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

Module

Пример:

>>> @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

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

Module

buffers(recurse=True) [исходный код]

Возвращает итератор по буферам модуля.

Параметры:

recurse (bool) – если True, возвращает буферы этого модуля и всех его подмодулей. В противном случае возвращает только буферы, являющиеся непосредственными членами этого модуля.

Возвращает значения:

torch.Tensor – буфер модуля

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

Iterator[Tensor]

Пример:

>>> for buf in model.buffers():
>>>     print(type(buf), buf.size())
<class 'torch.Tensor'> (20L,)
<class 'torch.Tensor'> (20L, 1L, 5L, 5L)
children() [исходный код]

Возвращает итератор по непосредственным дочерним модулям.

Возвращает значения:

Module – дочерний модуль

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

Iterator[Module]

compile(*args, **kwargs) [исходный код]

Компилирует метод forward этого модуля с помощью torch.compile().

Метод __call__ этого модуля компилируется, а все аргументы без изменений передаются в torch.compile().

Подробные сведения об аргументах этой функции см. в разделе torch.compile().

cpu() [исходный код]

Перемещает все параметры и буферы модели на CPU.

Примечание

Этот метод изменяет модуль на месте.

Возвращает:

self

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

Module

cuda(device=None) [исходный код]

Перемещает все параметры и буферы модели на GPU.

При этом связанные параметры и буферы также становятся другими объектами. Поэтому этот метод следует вызывать до создания оптимизатора, если модуль будет находиться на GPU во время оптимизации.

Примечание

Этот метод изменяет модуль на месте.

Параметры:

device (int, необязательно) – если указано, все параметры будут скопированы на это устройство

Возвращает:

self

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

Module

double() [исходный код]

Преобразует все параметры с плавающей точкой и буферы к типу данных double.

Примечание

Этот метод изменяет модуль на месте.

Возвращает:

self

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

Module

eval() [исходный код]

Переводит модуль в режим оценки.

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

Это эквивалентно вызову self.train(False).

Сравнение .eval() с несколькими похожими механизмами, которые можно с ним перепутать, см. в разделе Локальное отключение вычисления градиентов.

Возвращает:

self

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

Module

extra_repr() [исходный код]

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

Чтобы выводить дополнительную информацию в собственном формате, переопределите этот метод в своих модулях. Допускаются как однострочные, так и многострочные строки.

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

str

float() [исходный код]

Преобразует все параметры с плавающей точкой и буферы к типу данных float.

Примечание

Этот метод изменяет модуль на месте.

Возвращает:

self

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

Module

forward(*args, **kwargs) [исходный код]

Прямой проход.

get_buffer(target) [исходный код]

Возвращает буфер, указанный в target, если он существует; в противном случае вызывает ошибку.

Более подробное описание функциональности этого метода и правильного указания target см. в документации к get_submodule.

Параметры:

target (str) – полное строковое имя буфера, который необходимо найти. (См. get_submodule, где описано, как указать полную строку.)

Возвращает:

Буфер, указанный в target

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

torch.Tensor

Вызывает исключение:

AttributeError – если целевая строка указывает на недопустимый путь или разрешается в объект, который не является буфером

get_extra_state() [исходный код]

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

Если вашему модулю необходимо хранить дополнительные данные состояния, реализуйте этот метод и соответствующий ему метод set_extra_state(). Эта функция вызывается при создании state_dict() модуля.

Обратите внимание: дополнительные данные состояния должны поддерживать сериализацию с помощью pickle, чтобы сериализация state_dict работала корректно. Гарантии обратной совместимости предоставляются только для сериализации тензоров; изменение сериализованной формы pickle других объектов может нарушить обратную совместимость.

Возвращает:

Любые дополнительные данные состояния, сохраняемые в state_dict модуля

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

object

get_parameter(target) [исходный код]

Возвращает параметр, указанный в target, если он существует; в противном случае вызывает ошибку.

Более подробное описание функциональности этого метода и правильного указания target см. в документации к get_submodule.

Параметры:

target (str) – полное строковое имя параметра, который необходимо найти. (См. get_submodule, где описано, как указать полную строку.)

Возвращает:

Параметр, указанный в target

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

torch.nn.Parameter

Вызывает исключение:

AttributeError – если целевая строка указывает на недопустимый путь или разрешается в объект, который не является nn.Parameter

get_submodule(target) [исходный код]

Возвращает подмодуль, указанный в target, если он существует; в противном случае вызывает ошибку.

Например, предположим, что у вас есть nn.Module A следующей структуры:

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.Module A. 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

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

torch.nn.Module

Вызывает исключение:

AttributeError – если на любом этапе пути, заданного целевой строкой, под путь разрешается в несуществующее имя атрибута или в объект, который не является экземпляром nn.Module.

half() [исходный код]

Преобразует все параметры с плавающей точкой и буферы к типу данных half.

Примечание

Этот метод изменяет модуль на месте.

Возвращает:

self

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

Module

ipu(device=None) [исходный код]

Перемещает все параметры и буферы модели на IPU.

При этом связанные параметры и буферы также становятся другими объектами. Поэтому этот метод следует вызывать до создания оптимизатора, если модуль будет находиться на IPU во время оптимизации.

Примечание

Этот метод изменяет модуль на месте.

Параметры:

device (int, необязательно) – если указано, все параметры будут скопированы на это устройство

Возвращает:

self

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

Module

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 – модуль сети

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

Iterator[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 во время оптимизации.

Примечание

Этот метод изменяет модуль на месте.

Параметры:

device (int, optional) – если указан, все параметры будут скопированы на это устройство

Возвращает:

self

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

Module

named_buffers(prefix='', recurse=True, remove_duplicate=True) [исходный код]

Вернуть итератор по буферам модуля, выдавая как имя буфера, так и сам буфер.

Параметры:
  • prefix (str) – префикс, добавляемый ко всем именам буферов.
  • recurse (bool, optional) – если True, выдаёт буферы этого модуля и всех его подмодулей. В противном случае выдаёт только буферы, непосредственно принадлежащие этому модулю. По умолчанию True.
  • remove_duplicate (bool, optional) – удалять ли повторяющиеся буферы из результата. По умолчанию True.
Выдаёт:

(str, torch.Tensor) – кортеж, содержащий имя и буфер

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

Iterator[tuple[str, Tensor]]

Пример:

>>> for name, buf in self.named_buffers():
>>>     if name in ['running_var']:
>>>         print(buf.size())
named_children() [исходный код]

Вернуть итератор по непосредственным дочерним модулям, выдавая как имя модуля, так и сам модуль.

Выдаёт:

(str, Module) – кортеж, содержащий имя и дочерний модуль

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

Iterator[tuple[str, Module]]

Пример:

>>> for name, module in model.named_children():
>>>     if name in ['conv4', 'conv5']:
>>>         print(module)
named_modules(memo=None, prefix='', remove_duplicate=True) [исходный код]

Вернуть итератор по всем модулям сети, выдавая как имя модуля, так и сам модуль.

Параметры:
  • memo (set[Module] | None) – хранилище для множества модулей, уже добавленных в результат
  • prefix (str) – префикс, который будет добавлен к имени модуля
  • remove_duplicate (bool) – удалять ли повторяющиеся экземпляры модулей из результата
Выдаёт:

(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) – кортеж, содержащий имя и параметр

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

Iterator[tuple[str, Parameter]]

Пример:

>>> for name, param in self.named_parameters():
>>>     if name in ['bias']:
>>>         print(param.size())
parameters(recurse=True) [исходный код]

Вернуть итератор по параметрам модуля.

Точный порядок возвращаемых параметров не определён, однако повторные вызовы метода parameters() для неизменённого модуля возвращают параметры в одном и том же порядке.

Обычно этот итератор передают оптимизатору.

Параметры:

recurse (bool) – если True, выдаёт параметры этого модуля и всех его подмодулей. В противном случае выдаёт только параметры, непосредственно принадлежащие этому модулю.

Выдаёт:

Parameter – параметр модуля

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

Iterator[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_mean BatchNorm не является параметром, но входит в состояние модуля. По умолчанию буферы являются постоянными и сохраняются вместе с параметрами. Это поведение можно изменить, задав для 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) [исходный код]

Зарегистрировать хук обратного прохода для модуля.

Хук будет вызываться каждый раз при вычислении градиентов относительно модуля; правила его срабатывания следующие:

  1. Обычно хук срабатывает при вычислении градиентов относительно входных данных модуля.
  2. Если ни для одного из входных данных модуля не требуются градиенты, хук сработает при вычислении градиентов относительно выходных данных модуля.
  3. Если ни для одного из выходных данных модуля не требуются градиенты, хуки не сработают.

Хук должен иметь следующую сигнатуру:

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_() с несколькими похожими механизмами, которые можно с ним спутать.

Параметры:

requires_grad (bool) – должна ли autograd записывать операции с параметрами этого модуля. Значение по умолчанию: True.

Возвращает:

self

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

Module

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.Module A следующего вида:

A(
    (net_b): Module(
        (net_c): Module(
            (conv): Conv2d(3, 3, 3)
        )
        (linear): Linear(3, 3)
    )
)

(На диаграмме показан nn.Module A. 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.
share_memory() [исходный код]

См. 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.
Возвращает:

словарь, содержащий полное состояние модуля

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

dict

Пример:

>>> 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) – требуемый формат памяти для четырёхмерных параметров и буферов этого модуля (только именованный аргумент)
Возвращает:

self

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

Module

Примеры:

>>> 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) – нужно ли рекурсивно перемещать параметры и буферы подмодулей на указанное устройство.
Возвращает:

self

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

Module

train(mode=True) [исходный код]

Перевести модуль в режим обучения.

Это влияет только на некоторые модули. Подробные сведения о поведении отдельных модулей в режиме обучения/оценки см. в их документации, в частности о том, затрагиваются ли они, например, Dropout, BatchNorm и т. д.

Параметры:

mode (bool) – следует ли включить режим обучения (True) или режим оценки (False). Значение по умолчанию: True.

Возвращает:

self

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

Module

type(dst_type) [исходный код]

Привести все параметры и буферы к типу dst_type.

Примечание

Этот метод изменяет модуль на месте.

Параметры:

dst_type (type or string) – требуемый тип

Возвращает:

self

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

Module

update_parameters(model) [исходный код]

Обновить параметры модели.

xpu(device=None) [исходный код]

Переместить все параметры и буферы модели на XPU.

При этом связанные параметры и буферы также становятся другими объектами. Поэтому, если модуль будет находиться на XPU во время оптимизации, этот метод следует вызывать до создания оптимизатора.

Примечание

Этот метод изменяет модуль на месте.

Параметры:

device (int, optional) – если указан, все параметры будут скопированы на это устройство

Возвращает:

self

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

Module

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

Spec-Zone.ru

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