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)[исходный код] -
Добавляет дочерний модуль в текущий модуль.
К модулю можно обратиться как к атрибуту, используя указанное имя.
-
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(*input)[исходный код] -
Определяет вычисления, выполняемые при каждом вызове.
Этот метод должны переопределять все подклассы.
Примечание
Хотя алгоритм прямого прохода необходимо определить в этой функции, вместо неё следует вызывать экземпляр
Module, поскольку он запускает зарегистрированные хуки, тогда как прямой вызов этой функции незаметно их игнорирует.
-
get_buffer(target)[исходный код] -
Возвращает буфер, указанный в
target, если он существует; в противном случае вызывает ошибку.Более подробное описание функциональности этого метода и правильного задания
targetсм. в документацииget_submodule.- Параметры:
-
target (str) – полное строковое имя буфера для поиска. (О том, как задавать полное строковое имя, см.
get_submodule.) - Возвращает:
-
Буфер, на который ссылается
target - Тип возвращаемого значения:
- Вызывает исключение:
-
AttributeError – если строка target указывает на недопустимый путь или разрешается в объект, не являющийся буфером
-
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 – если строка target указывает на недопустимый путь или разрешается в объект, не являющийся
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 – если на любом участке пути, заданного строкой target, подстрока пути разрешается в несуществующее имя атрибута или объект, не являющийся экземпляром
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, необязательный) – следует ли строго проверять совпадение ключей в
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 – модуль сети
- Тип возвращаемого значения:
Примечание
По умолчанию повторяющиеся модули возвращаются только один раз. В следующем примере
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_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)[исходный код] -
Регистрирует хук обратного прохода для модуля.
Хук будет вызываться каждый раз при вычислении градиентов относительно модуля; правила его вызова следующие:
- Обычно хук вызывается при вычислении градиентов относительно входных данных модуля.
- Если для входных данных модуля не требуются градиенты, хук будет вызван при вычислении градиентов относительно выходных данных модуля.
- Если для выходных данных модуля не требуются градиенты, хуки вызваны не будут.
Сигнатура хука должна быть следующей:
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_()с несколькими похожими механизмами, которые можно с ним перепутать.
-
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.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(*, 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.
-
destination (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-параметров и буферов этого модуля (только именованный аргумент)
-
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.Примечание
Этот метод изменяет модуль на месте.
-
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.nn.Module.html