Модуль
-
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] -
Добавляет дочерний модуль в текущий модуль.
Модуль можно получить как атрибут с помощью заданного имени.
-
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] -
Возвращает итератор по непосредственным дочерним модулям.
-
compile(*args, **kwargs)[source] -
Компилирует прямой проход этого модуля с использованием
torch.compile().Метод
__call__этого модуля компилируется, и все аргументы передаются как есть вtorch.compile().См.
torch.compile()для получения подробной информации об аргументах этой функции.
-
cpu()[source] -
Перемещает все параметры и буферы модели на процессор.
Примечание
Этот метод изменяет модуль на месте.
- Возвращает
-
self
- Тип возвращаемого значения
-
cuda(device=None)[source] -
Перемещает все параметры и буферы модели на графический процессор.
Это также создаёт новые объекты для связанных параметров и буферов. Поэтому его следует вызывать перед построением оптимизатора, если модуль будет находиться на графическом процессоре во время оптимизации.
Примечание
Этот метод изменяет модуль на месте.
-
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()модуля.Обратите внимание, что дополнительное состояние должно быть сериализуемым для обеспечения работы сериализации 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, 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] -
Возвращает итератор по всем модулям в сети.
Примечание
Дублируемые модули возвращаются только один раз. В следующем примере,
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) – Кортеж, содержащий имя и буфер
- Тип возвращаемого значения
Пример:
>>> for name, buf in self.named_buffers(): >>> if name in ['running_var']: >>> print(buf.size())
-
named_children()[source] -
Возвращает итератор по непосредственным дочерним модулям, вынося как имя модуля, так и сам модуль.
- Возвращает
-
(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] -
Возвращает итератор по всем модулям в сети, вынося как имя модуля, так и сам модуль.
- Параметры
- Возвращает
-
(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– словарьliststr, содержащий пропущенные ключи, аunexpected_keys– словарьliststr, содержащий неожиданные ключи.Предоставленный словарь 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_()и несколькими похожими механизмами, которые могут быть с ним перепутаны.
-
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, необязательно) – Если указано, состояние модуля будет обновлено в словаре, и возвращается тот же объект. В противном случае будет создан и возвращен новый
OrderedDict. По умолчанию:None. -
prefix (str, необязательно) – префикс, добавляемый к именам параметров и буферов для составления ключей в state_dict. По умолчанию:
''. -
keep_vars (bool, необязательно) – по умолчанию
Tensors, возвращаемые в state_dict, отсоединяются от автограда. Если установлено в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, recurse=True)[source] -
Перемещает параметры и буферы на указанное устройство без копирования хранения.
- Параметры
-
-
device (
torch.device) – Желаемое устройство параметров и буферов в этом модуле. - recurse (bool) – Нужно ли рекурсивно перемещать параметры и буферы подмодулей на указанное устройство.
-
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=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