Spec-Zone.ru › PyTorch 2

Модуль

class torch.nn.Module(*args, **kwargs) [source]

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

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

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

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

class Model(nn.Module):
    def __init__(self):
        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() и т.д.

Примечание

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

Переменные

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

add_module(name, module) [source]

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

Модуль можно получить как атрибут с помощью заданного имени.

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

Применяет fn рекурсивно к каждому подмодулю (как возвращается .children()) а также к self. Типичное использование включает инициализацию параметров модели (см. также torch.nn.init).

Параметры

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

Возвращает

self

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

Module

Пример:

>>> @torch.no_grad()
>>> def init_weights(m):
>>>     print(m)
>>>     if type(m) == 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() [source]

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

Примечание

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

Возвращает

self

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

Module

buffers(recurse=True) [source]

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

Параметры

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() [source]

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

Возвращает

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

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

Iterator[Module]

compile(*args, **kwargs) [source]

Компилирует прямой проход этого модуля с использованием torch.compile().

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

См. torch.compile() для получения подробной информации об аргументах этой функции.

cpu() [source]

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

Примечание

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

Возвращает

self

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

Module

cuda(device=None) [source]

Перемещает все параметры и буферы модели на графический процессор.

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

Примечание

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

Параметры

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

Возвращает

self

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

Module

double() [source]

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

Примечание

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

Возвращает

self

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

Module

eval() [source]

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

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

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

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

Возвращает

self

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

Module

extra_repr() [source]

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

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

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

str

float() [source]

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

Примечание

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

Возвращает

self

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

Модуль

forward(*input)

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

Должно быть переопределено всеми подклассами.

Примечание

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

get_buffer(target) [source]

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

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

Параметры

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

Возвращает

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

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

torch.Tensor

Исключения

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

get_extra_state() [source]

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

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

Возвращает

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

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

object

get_parameter(target) [source]

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

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

Параметры

target (str) – Полное квалифицированное имя строки параметра для поиска. (См. get_submodule для указания полной квалифицированной строки.)

Возвращает

Параметр, на который ссылается target

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

torch.nn.Parameter

Исключения

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

get_submodule(target) [source]

Возвращает подмодуль, заданный 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() [source]

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

Примечание

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

Возвращает

self

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

Модуль

ipu(device=None) [source]

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

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

Примечание

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

Параметры

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

Возвращает

self

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

Модуль

load_state_dict(state_dict, strict=True, assign=False) [source]

Копирует параметры и буферы из state_dict в данный модуль и его потомки. Если strict равно True, то ключи словаря state_dict должны точно совпадать с ключами, возвращаемыми функцией state_dict() данного модуля.

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

Если assign равно True, оптимизатор необходимо создавать после вызова load_state_dict.

Параметры
  • state_dict (dict) – словарь, содержащий параметры и постоянные буферы.
  • strict (bool, необязательно) – строго ли следовать тому, что ключи в state_dict должны совпадать с ключами, возвращаемыми функцией state_dict() данного модуля. По умолчанию: True
  • assign (bool, необязательно) – присваивать ли элементы словаря состояния соответствующим ключам модуля вместо копирования их напрямую в текущие параметры и буферы модуля. При False, свойства тензоров в текущем модуле сохраняются, а при True, свойства тензоров в словаре состояния сохраняются. По умолчанию: False
Возвращает
  • missing_keys – список строк, содержащий пропущенные ключи
  • unexpected_keys – список строк, содержащий неожиданные ключи
Тип возвращаемого значения

NamedTuple со missing_keys и unexpected_keys полями

Примечание

Если параметр или буфер зарегистрирован как None и соответствующий ключ существует в state_dict, load_state_dict() вызовет RuntimeError.

modules() [source]

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

Возвращает

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)
named_buffers(prefix='', recurse=True, remove_duplicate=True) [source]

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

Параметры
  • prefix (str) – префикс, который будет добавлен ко всем именам буферов.
  • recurse (bool, необязательно) – если True, то возвращает буферы этого модуля и всех подмодулей. В противном случае, возвращает только буферы, которые являются прямыми членами этого модуля. По умолчанию True.
  • remove_duplicate (bool, необязательно) – удалять ли дублирующиеся буферы в результате. По умолчанию 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() [source]

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

Возвращает

(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) [source]

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

Параметры
  • memo (Optional[Set[Module]]) – запоминание, чтобы сохранить множество модулей, уже добавленных в результат
  • 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) [source]

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

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

(строка, Параметр) – Кортеж, содержащий имя и параметр

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

Итератор[Кортеж[строка, Параметр]]

Пример:

>>> for name, param in self.named_parameters():
>>>     if name in ['bias']:
>>>         print(param.size())
parameters(recurse=True) [source]

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

Это обычно передается в оптимизатор.

Параметры

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

Возвращаемые значения

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

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

Итератор[Параметр]

Пример:

>>> 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) [source]

Регистрирует обратный хук для модуля.

Эта функция устарела в пользу register_full_backward_hook() и поведение этой функции изменится в будущих версиях.

Возвращаемые значения

дескриптор, который можно использовать для удаления добавленного хука, вызвав handle.remove()

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

torch.utils.hooks.RemovableHandle

register_buffer(name, tensor, persistent=True) [source]

Добавляет буфер к модулю.

Это обычно используется для регистрации буфера, который не должен рассматриваться как параметр модели. Например, running_mean в BatchNorm не является параметром, но является частью состояния модуля. Буферы по умолчанию постоянны и будут сохранены вместе с параметрами. Это поведение можно изменить, установив persistent в False. Единственное различие между постоянным буфером и непостоянным буфером заключается в том, что последний не будет частью state_dict этого модуля.

К буферам можно получить доступ как к атрибутам с помощью заданных имен.

Параметры
  • name (строка) – имя буфера. К буферу можно получить доступ из этого модуля с помощью заданного имени
  • tensor (Тензор или None) – регистрируемый буфер. Если None, операции, выполняемые над буферами, такие как cuda, игнорируются. Если None, буфер не включается в state_dict модуля.
  • persistent (булево значение) – является ли буфер частью state_dict этого модуля.

Пример:

>>> self.register_buffer('running_mean', torch.zeros(num_features))
register_forward_hook(hook, *, prepend=False, with_kwargs=False, always_call=False) [source]

Регистрирует прямой хук для модуля.

Хук вызывается каждый раз после того, как forward() вычислил выход.

Если with_kwargs равно False или не указано, вход содержит только позиционные аргументы, предоставленные модулю. Ключевые аргументы не будут переданы в хуки, а только в forward. Хук может изменить выход. Он может изменять вход in-place, но это не повлияет на прямой вызов, так как это происходит после вызова forward(). Хук должен иметь следующий синтаксис:

hook(module, args, output) -> None or modified output

Если with_kwargs равно True, прямой хук будет получать kwargs, предоставленные функции forward, и ожидается, что он вернёт выход, возможно, изменённый. Хук должен иметь следующий синтаксис:

hook(module, args, kwargs, output) -> None or modified output
Параметры
  • hook (Вызываемая функция) – Пользовательский хук для регистрации.
  • prepend (булево значение) – Если True, предоставленный hook будет запущен до всех существующих forward хуков в этом torch.nn.modules.Module. В противном случае предоставленный hook будет запущен после всех существующих forward хуков в этом torch.nn.modules.Module. Обратите внимание, что глобальные forward хуки, зарегистрированные с помощью register_module_forward_hook(), будут срабатывать до всех хуков, зарегистрированных этим методом. По умолчанию: False
  • with_kwargs (булево значение) – Если True, hook получит kwargs, переданные функции forward. По умолчанию: False
  • always_call (булево значение) – Если True хук hook будет выполнен независимо от того, возникает ли исключение при вызове Module. По умолчанию: False
Возвращаемые значения

дескриптор, который можно использовать для удаления добавленного хука, вызвав handle.remove()

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

torch.utils.hooks.RemovableHandle

register_forward_pre_hook(hook, *, prepend=False, with_kwargs=False) [source]

Регистрирует предварительный хук для метода forward() модуля.

Хук будет вызываться каждый раз перед вызовом forward().

Если with_kwargs ложно или не указано, вход содержит только позиционные аргументы, переданные модулю. Ключевые аргументы не будут переданы хукам и только модулю. Хук может изменить вход. Пользователь может вернуть кортеж или единственное изменённое значение в хуке. Мы обернём значение в кортеж, если возвращено единственное значение (если это не кортеж). Хук должен иметь следующий вид:

hook(module, args) -> None or modified input

Если with_kwargs истинно, предварительный хук метода forward() будет принимать kwargs, переданные функции forward. И если хук изменяет вход, должны быть возвращены как args, так и kwargs. Хук должен иметь следующий вид:

hook(module, args, kwargs) -> None or a tuple of modified input and kwargs
Параметры
  • hook (Callable) – Определяемый пользователем хук для регистрации.
  • prepend (bool) – Если истинно, предоставленный hook будет вызван до всех существующих forward_pre хуков в этом torch.nn.modules.Module. В противном случае, предоставленный hook будет вызван после всех существующих forward_pre хуков в этом torch.nn.modules.Module. Обратите внимание, что глобальные forward_pre хуки, зарегистрированные с register_module_forward_pre_hook(), будут вызваны до всех хуков, зарегистрированных этим методом. Значение по умолчанию: False
  • with_kwargs (bool) – Если истинно, hook будет принимать kwargs, переданные функции forward. Значение по умолчанию: False
Возвращаемое значение

дескриптор, который можно использовать для удаления добавленного хука, вызвав handle.remove()

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

torch.utils.hooks.RemovableHandle

register_full_backward_hook(hook, prepend=False) [source]

Регистрирует хук обратного распространения ошибки для модуля.

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

hook(module, grad_input, grad_output) -> tuple(Tensor) or None

grad_input и grad_output – кортежи, содержащие градиенты по отношению к входам и выходам соответственно. Хук не должен изменять свои аргументы, но он может возвращать новый градиент по отношению к входу, который будет использоваться вместо grad_input в последующих вычислениях. grad_input будет соответствовать только входам, переданным в качестве позиционных аргументов, и все аргументы kwarg игнорируются. Элементы в grad_input и grad_output будут None для всех аргументов, не являющихся тензорами.

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

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

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

Параметры
  • hook (Callable) – Хук, определённый пользователем.
  • prepend (bool) – Если истинно, предоставленный hook будет вызван до всех существующих backward хуков в этом torch.nn.modules.Module. В противном случае, предоставленный hook будет вызван после всех существующих backward хуков в этом torch.nn.modules.Module. Обратите внимание, что глобальные backward хуки, зарегистрированные с register_module_full_backward_hook(), будут вызваны до всех хуков, зарегистрированных этим методом.
Возвращаемое значение

дескриптор, который можно использовать для удаления добавленного хука, вызвав handle.remove()

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

torch.utils.hooks.RemovableHandle

register_full_backward_pre_hook(hook, prepend=False) [source]

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

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

hook(module, grad_output) -> tuple[Tensor] or None

grad_output – кортеж. Хук не должен изменять свои аргументы, но он может возвращать новый градиент по отношению к выходу, который будет использоваться вместо grad_output в последующих вычислениях. Элементы в grad_output будут None для всех аргументов, не являющихся тензорами.

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

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

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

Параметры
  • hook (Callable) – Хук, определённый пользователем.
  • prepend (bool) – Если истинно, предоставленный hook будет вызван до всех существующих backward_pre хуков в этом torch.nn.modules.Module. В противном случае, предоставленный hook будет вызван после всех существующих backward_pre хуков в этом torch.nn.modules.Module. Обратите внимание, что глобальные backward_pre хуки, зарегистрированные с register_module_full_backward_pre_hook(), будут вызваны до всех хуков, зарегистрированных этим методом.
Возвращаемое значение

дескриптор, который можно использовать для удаления добавленного хука, вызвав handle.remove()

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

torch.utils.hooks.RemovableHandle

register_load_state_dict_post_hook(hook) [source]

Регистрирует пост-хук, который выполняется после вызова метода загрузки состояния модуля load_state_dict.

Он должен иметь следующий вид::

hook(module, incompatible_keys) -> None

Аргумент module – текущий модуль, на котором зарегистрирован этот хук, а аргумент incompatible_keys – словарь, содержащий атрибуты 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_module(name, module) [source]

Псевдоним для add_module().

register_parameter(name, param) [source]

Добавляет параметр в модуль.

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

Параметры
  • name (str) – имя параметра. К параметру можно получить доступ из этого модуля по данному имени
  • param (Parameter или None) – параметр, который нужно добавить в модуль. Если None, операции, выполняемые над параметрами, такие как cuda, игнорируются. Если None, параметр не включается в state_dict модуля.
register_state_dict_pre_hook(hook) [source]

Эти хуки вызываются с аргументами: self, prefix, и keep_vars перед вызовом state_dict на self. Зарегистрированные хуки могут использоваться для выполнения предобработки перед вызовом state_dict.

requires_grad_(requires_grad=True) [source]

Изменяет, следует ли автограду записывать операции над параметрами в этом модуле.

Этот метод устанавливает атрибуты requires_grad параметров на месте.

Этот метод полезен для заморозки части модуля для дообучения или обучения частей модели по отдельности (например, обучение GAN).

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

Параметры

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

Возвращает

self

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

Модуль

set_extra_state(state) [source]

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

Параметры

state (dict) – Дополнительное состояние из state_dict

share_memory() [source]

См. torch.Tensor.share_memory_()

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

T

state_dict(*, destination: T_destination, prefix: str = '', keep_vars: bool = False) → T_destination [source]
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 s, возвращаемые в state_dict, отсоединяются от автограда. Если установлено в True, отсоединение не будет выполняться. По умолчанию: False.
Возвращает

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

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

dict

Пример:

>>> module.state_dict().keys()
['bias', 'weight']
to(device: Optional[Union[int, device]] = ..., dtype: Optional[Union[dtype, str]] = ..., non_blocking: bool = ...) → T [source]
to(dtype:Union[dtype,str], non_blocking:bool=...) → T
to(tensor:Tensor, non_blocking:bool=...) → T

Перемещает и/или преобразует параметры и буферы.

Можно вызвать как

to(device=None, dtype=None, non_blocking=False)[source]
to(dtype, non_blocking=False)[source]
to(tensor, non_blocking=False)[source]
to(memory_format=torch.channels_last)[source]

Её сигнатура похожа на 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

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

Модуль

Примеры:

>>> 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) [source]

Перемещает параметры и буферы на указанное устройство без копирования хранения.

Параметры
  • device (torch.device) – Желаемое устройство параметров и буферов в этом модуле.
  • recurse (bool) – Нужно ли рекурсивно перемещать параметры и буферы подмодулей на указанное устройство.
Возвращает

self

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

Модуль

train(mode=True) [source]

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

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

Параметры

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

Возвращает

self

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

Модуль

type(dst_type) [source]

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

Примечание

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

Параметры

dst_type (type или строка) – желаемый тип

Возвращает

self

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

Модуль

xpu(device=None) [source]

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

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

Примечание

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

Параметры

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

Возвращает

self

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

Модуль

zero_grad(set_to_none=True) [source]

Сбрасывает градиенты всех параметров модели. Смотрите аналогичную функцию в torch.optim.Optimizer для более подробной информации.

Параметры

set_to_none (bool) – вместо сброса в ноль, установить градиенты в None. См. torch.optim.Optimizer.zero_grad() для подробностей.

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.nn.Module.html

Spec-Zone.ru

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