Spec-Zone.ru › PyTorch 1

Модуль

class torch.nn.Module [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]

cpu() [source]

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

Примечание

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

Возвращает:

self

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

Module

cuda(device=None) [source]

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

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

Примечание

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

Параметры:

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

Возвращает:

self

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

Module

double() [source]

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

Примечание

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

Возвращает:

self

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

Module

eval() [source]

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

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

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

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

Возвращает:

self

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

Module

extra_repr() [source]

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

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

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

str

END_OF_DOCUMENT_MARKER
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() модуля.

Обратите внимание, что дополнительное состояние должно быть сериализуемым с помощью pickle, чтобы гарантировать правильную сериализацию 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

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

Модуль

END_OF_DOCUMENT_MARKER
load_state_dict(state_dict, strict=True) [source]

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

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

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

Примечание

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

modules() [source]

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

Возвращает:

Модуль – модуль в сети

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

Iterator[Модуль]

Примечание

Дублируемые модули возвращаются только один раз. В следующем примере, 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) [source]

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

Параметры:
  • prefix (str) – префикс, добавляемый ко всем именам буферов.
  • recurse (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, Модуль) – кортеж, содержащий имя и дочерний модуль

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

Iterator[Tuple[str, Модуль]]

Пример:

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

(str, Модуль) – кортеж из имени и модуля

Примечание

Дублируемые модули возвращаются только один раз. В следующем примере, 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) [source]

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

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

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

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

Iterator[Tuple[str, Параметр]]

Пример:

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

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

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

Параметры:

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

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

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

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

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

Возвращает:

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

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

torch.utils.hooks.RemovableHandle

register_forward_pre_hook(hook) [source]

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

Хук будет вызываться каждый раз перед вызовом forward(). Он должен иметь следующий вид:

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

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

Возвращает:

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

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

torch.utils.hooks.RemovableHandle

register_full_backward_hook(hook) [source]

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

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

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

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

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

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

Изменение входных или выходных данных in-place запрещено при использовании обратных хуков и вызовет ошибку.

Возвращает:

указатель, который можно использовать для удаления добавленного хука, вызвав 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, содержащий нежелательные ключи.

Данный кортеж incompatible_keys можно изменять in-place при необходимости.

Обратите внимание, что проверки, выполняемые при вызове 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 or None) – параметр, который необходимо добавить в модуль. Если None, тогда операции, выполняемые над параметрами, такие как cuda, игнорируются. Если None, параметр не включается в state_dict модуля.
requires_grad_(requires_grad=True) [source]

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

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

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

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

Параметры:

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

Возвращает:

self

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

Module

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, необязательно) – Если указано, состояние модуля будет обновлено в словаре, и возвращается тот же объект. В противном случае будет создан и возвращён новый словарь. По умолчанию: None.
  • prefix (str, необязательно) – префикс, добавляемый к именам параметров и буферов для составления ключей в state_dict. По умолчанию: ''.
  • keep_vars (bool, необязательно) – по умолчанию, Tensor , возвращаемые в словаре state_dict, отделены от autograd. Если установлено значение 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) [source]

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

Параметры:

device (torch.device) – Желаемое устройство для параметров и буферов в этом модуле.

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

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=False) [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/1.13/generated/torch.nn.Module.html

Spec-Zone.ru

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