Spec-Zone.ru › PyTorch 2.14

Module

class torch.nn.Module(*args, **kwargs) [исходный код]

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

Ваши модели также должны наследовать этот класс.

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

import torch.nn as nn
import torch.nn.functional as F


class Model(nn.Module):
    def __init__(self) -> None:
        super().__init__()
        self.conv1 = nn.Conv2d(1, 20, 5)
        self.conv2 = nn.Conv2d(20, 20, 5)

    def forward(self, x):
        x = F.relu(self.conv1(x))
        return F.relu(self.conv2(x))

Подмодули, назначенные таким образом, будут зарегистрированы, а их параметры также будут преобразованы при вызове to() и т. д.

Примечание

Как показано в примере выше, вызов __init__() родительского класса необходимо выполнить до назначения атрибута дочернему классу.

Переменные:

training (bool) – логическое значение, указывающее, находится ли этот модуль в режиме обучения или оценки.

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

Определяет вычисления, выполняемые при каждом вызове.

Этот метод должны переопределять все подклассы.

Примечание

Хотя алгоритм прямого прохода необходимо определить в этой функции, вместо неё следует вызывать экземпляр Module, поскольку он запускает зарегистрированные хуки, тогда как прямой вызов этой функции незаметно их игнорирует.

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

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

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

Параметры:

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

Возвращает:

Буфер, на который ссылается target

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

torch.Tensor

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

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

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 – если строка target указывает на недопустимый путь или разрешается в объект, не являющийся 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 – если на любом участке пути, заданного строкой target, подстрока пути разрешается в несуществующее имя атрибута или объект, не являющийся экземпляром 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, необязательный) – следует ли строго проверять совпадение ключей в state_dict с ключами, возвращаемыми функцией state_dict() этого модуля. По умолчанию: True
  • assign (bool, необязательный) – если задано значение 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, необязательный) – если задан, все параметры будут скопированы на это устройство

Возвращает:

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 для всех аргументов, не являющихся Tensor.

По техническим причинам при применении этого хука к Module его функция прямого прохода получает представление каждого переданного модулю Tensor. Аналогично вызывающая сторона получает представление каждого Tensor, возвращённого функцией прямого прохода модуля.

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

При использовании хуков обратного прохода запрещено изменять входные или выходные данные на месте; такая попытка приведёт к ошибке.

Параметры:
  • 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 для всех аргументов, не являющихся Tensor.

По техническим причинам при применении этого хука к Module его функция прямого прохода получает представление каждого переданного модулю Tensor. Аналогично вызывающая сторона получает представление каждого Tensor, возвращённого функцией прямого прохода модуля.

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

При использовании хуков обратного прохода запрещено изменять входные данные на месте; такая попытка приведёт к ошибке.

Параметры:
  • hook (Callable) – определённый пользователем хук для регистрации.
  • prepend (bool) – если true, предоставленный 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. Если вам нужно хранить дополнительные данные состояния в state_dict модуля, реализуйте эту функцию и соответствующую ей функцию get_extra_state().

Параметры:

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(*, destination: T_destination, prefix: str = '', keep_vars: bool = False) → T_destination [исходный код]
state_dict(*, prefix:str='', keep_vars:bool=False) → dict[str,Any]

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

В него входят параметры и постоянные буферы (например, скользящие средние). Ключи соответствуют именам параметров и буферов. Параметры и буферы, для которых задано None, не включаются.

Примечание

Возвращаемый объект является поверхностной копией. Он содержит ссылки на параметры и буферы модуля.

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

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

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

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

Параметры:
  • destination (dict, необязательный) – Если указан, состояние модуля будет записано в этот словарь, и будет возвращён тот же объект. В противном случае будет создан и возвращён OrderedDict. Значение по умолчанию: None.
  • prefix (str, необязательный) – префикс, добавляемый к именам параметров и буферов для формирования ключей в state_dict. Значение по умолчанию: ''.
  • keep_vars (bool, необязательный) – по умолчанию Tensor, возвращаемые в словаре состояния, отсоединены от autograd. Если задано True, отсоединение не выполняется. Значение по умолчанию: False.
Возвращает:

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

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

dict

Пример:

>>> module.state_dict().keys()
['bias', 'weight']
to(device: str | device | int | None = ..., dtype: dtype | None = ..., non_blocking: bool = ...) → Self [исходный код]
to(dtype:dtype, non_blocking:bool=...) → Self
to(tensor:Tensor, non_blocking:bool=...) → Self
to(memory_format:memory_format) → Self

Переместить и/или преобразовать тип параметров и буферов.

Этот метод можно вызывать следующим образом:

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

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 или строка) – целевой тип

Возвращает:

self

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

Module

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

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

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

Примечание

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

Параметры:

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

Возвращает:

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.nn.Module.html

Spec-Zone.ru

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