Модуль
-
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] -
Добавляет дочерний модуль в текущий модуль.
Модуль можно получить в качестве атрибута с заданным именем.
-
apply(fn)[source] -
Применяет
fnрекурсивно ко всем подмодулям (как возвращает.children()) а также к self. Типичное использование включает инициализацию параметров модели (см. также torch.nn.init).- Параметры:
-
fn (
Module-> None) – функция, которая должна быть применена к каждому подмодулю - Возвращает:
-
self
- Тип возвращаемого значения:
Пример:
>>> @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
- Тип возвращаемого значения:
-
buffers(recurse=True)[source] -
Возвращает итератор по буферам модуля.
- Параметры:
-
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()[source] -
Возвращает итератор по непосредственным дочерним модулям.
-
cpu()[source] -
Перемещает все параметры и буферы модели на CPU.
Примечание
Этот метод изменяет модуль на месте.
- Возвращает:
-
self
- Тип возвращаемого значения:
-
cuda(device=None)[source] -
Перемещает все параметры и буферы модели на GPU.
Это также делает связанные параметры и буферы разными объектами. Поэтому его следует вызывать до построения оптимизатора, если модуль будет существовать на GPU во время оптимизации.
Примечание
Этот метод изменяет модуль на месте.
-
double()[source] -
Преобразует все параметры с плавающей запятой и буферы в тип данных
double.Примечание
Этот метод изменяет модуль на месте.
- Возвращает:
-
self
- Тип возвращаемого значения:
-
eval()[source] -
Устанавливает модуль в режим оценки.
Это оказывает влияние только на определенные модули. Подробности о поведении в режиме обучения/оценки для конкретных модулей см. в документации по соответствующим модулям, если они затронуты, например,
Dropout,BatchNorm, и т.д.Это эквивалентно
self.train(False).См. Локальное отключение вычисления градиента для сравнения между
.eval()и несколькими аналогичными механизмами, которые могут с ним путаться.- Возвращает:
-
self
- Тип возвращаемого значения:
-
extra_repr()[source] -
Установить дополнительное представление модуля
Чтобы вывести на печать пользовательскую дополнительную информацию, вы должны повторно реализовать этот метод в собственных модулях. Допускаются как однострочные, так и многострочные строки.
- Тип возвращаемого значения:
-
float()[source] -
Преобразует все параметры и буферы с плавающей точкой в тип данных
float.Примечание
Этот метод изменяет модуль на месте.
- Возвращаемое значение:
-
self
- Тип возвращаемого значения:
-
forward(*input) -
Определяет вычисление, выполняемое при каждом вызове.
Должен быть переопределен всеми подклассами.
Примечание
Хотя рецепт прямого прохода должен быть определён в этой функции, следует вызывать экземпляр
Moduleпосле этого, а не эту функцию, так как первый вариант обрабатывает запуск зарегистрированных хуков, а второй их молча игнорирует.
-
get_buffer(target)[source] -
Возвращает буфер, заданный
target, если он существует, иначе генерирует ошибку.См. строку документации для
get_submoduleдля более подробного объяснения функциональности этого метода, а также о том, как правильно указатьtarget.- Параметры:
-
target (str) – Полное квалифицированное строковое имя буфера для поиска. (См.
get_submoduleо том, как указать полное квалифицированное строковое имя.) - Возвращаемое значение:
-
Буфер, на который ссылается
target - Тип возвращаемого значения:
- Исключения:
-
AttributeError – Если целевая строка ссылается на неверный путь или указывает на что-то, что не является буфером
-
get_extra_state()[source] -
Возвращает дополнительное состояние для включения в state_dict модуля. Реализуйте это и соответствующий метод
set_extra_state()для вашего модуля, если вам нужно хранить дополнительное состояние. Эта функция вызывается при построенииstate_dict()модуля.Обратите внимание, что дополнительное состояние должно быть сериализуемым с помощью pickle, чтобы гарантировать правильную сериализацию state_dict. Мы гарантируем обратную совместимость только для сериализации тензоров; другие объекты могут нарушить обратную совместимость, если их сериализованный вид в формате pickle изменится.
- Возвращаемое значение:
-
Любое дополнительное состояние для хранения в state_dict модуля
- Тип возвращаемого значения:
-
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.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 – Если целевая строка ссылается на неверный путь или указывает на что-то, что не является
nn.Module
-
half()[source] -
Преобразует все параметры и буферы с плавающей точкой в тип данных
half.Примечание
Этот метод изменяет модуль на месте.
- Возвращаемое значение:
-
self
- Тип возвращаемого значения:
-
ipu(device=None)[source] -
Перемещает все параметры и буферы модели на IPU.
Это также создаёт разные объекты для связанных параметров и буферов. Поэтому его следует вызывать до построения оптимизатора, если модуль будет существовать на IPU во время оптимизации.
Примечание
Этот метод изменяет модуль на месте.
-
-
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] -
Возвращает итератор по всем модулям в сети.
Примечание
Дублируемые модули возвращаются только один раз. В следующем примере,
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] -
Возвращает итератор по буферам модуля, возвращая имя буфера и сам буфер.
- Параметры:
- Возвращает:
-
(str, torch.Tensor) – кортеж, содержащий имя и буфер
- Тип возвращаемого значения:
Пример:
>>> for name, buf in self.named_buffers(): >>> if name in ['running_var']: >>> print(buf.size())
-
named_children()[source] -
Возвращает итератор по непосредственным дочерним модулям, возвращая имя модуля и сам модуль.
- Возвращает:
-
(str, Модуль) – кортеж, содержащий имя и дочерний модуль
- Тип возвращаемого значения:
Пример:
>>> for name, module in model.named_children(): >>> if name in ['conv4', 'conv5']: >>> print(module)
-
named_modules(memo=None, prefix='', remove_duplicate=True)[source] -
Возвращает итератор по всем модулям в сети, возвращая имя модуля и сам модуль.
- Параметры:
- Возвращает:
-
(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] -
Возвращает итератор по параметрам модуля, возвращая имя параметра и сам параметр.
- Параметры:
- Возвращает:
-
(str, Параметр) – кортеж, содержащий имя и параметр
- Тип возвращаемого значения:
Пример:
>>> for name, param in self.named_parameters(): >>> if name in ['bias']: >>> print(param.size())
-
-
parameters(recurse=True)[source] -
Возвращает итератор по параметрам модуля.
Обычно передаётся в оптимизатор.
- Параметры:
-
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)[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_()и несколькими аналогичными механизмами, которые могут быть с ним перепутаны.
-
set_extra_state(state)[source] -
Эта функция вызывается из
load_state_dict()для обработки любого дополнительного состояния, найденного вstate_dict. Реализуйте эту функцию и соответствующуюget_extra_state()для своего модуля, если вам нужно хранить дополнительное состояние внутри егоstate_dict.- Параметры:
-
state (dict) – Дополнительное состояние из
state_dict
-
См.
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.
-
destination (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 параметров и буферов в этом модуле (только ключевой аргумент)
-
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)[source] -
Перемещает параметры и буферы на указанное устройство без копирования хранилища.
- Параметры:
-
device (
torch.device) – Желаемое устройство для параметров и буферов в этом модуле. - Возвращаемое значение:
-
self
- Тип возвращаемого значения:
-
train(mode=True)[source] -
Устанавливает модуль в режим обучения.
Это оказывает влияние только на определённые модули. См. документацию конкретных модулей для получения подробной информации о их поведении в режиме обучения/оценки, если они затронуты, например,
Dropout,BatchNorm, и т. д.
-
type(dst_type)[source] -
Преобразует все параметры и буферы к
dst_type.Примечание
Этот метод изменяет модуль на месте.
-
xpu(device=None)[source] -
Перемещает все параметры и буферы модели на XPU.
Это также создаёт отдельные объекты для связанных параметров и буферов. Поэтому его следует вызывать перед построением оптимизатора, если модуль будет жить на XPU во время оптимизации.
Примечание
Этот метод изменяет модуль на месте.
-
-
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