Spec-Zone.ru › PyTorch 1

МодульScriptModule

class torch.jit.ScriptModule [исходный код]

Обёртка вокруг C++ torch::jit::Module. ScriptModule содержат методы, атрибуты, параметры и константы. К ним можно обратиться так же, как и к обычным nn.Module.

add_module(name, module)

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

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

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

Применяет 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()

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

Примечание

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

Возвращает:

self

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

Модуль

buffers(recurse=True)

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

Параметры:

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

Возвращает:

torch.Tensor – буфер модуля

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

Итератор[Тензор]

Пример:

>>> for buf in model.buffers():
>>>     print(type(buf), buf.size())
<class 'torch.Tensor'> (20L,)
<class 'torch.Tensor'> (20L, 1L, 5L, 5L)
children()

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

Возвращает:

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

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

Итератор[Модуль]

property code

Возвращает красиво отформатированное представление (в виде корректного синтаксиса Python) внутреннего графа для метода forward. Подробнее см. Просмотр кода.

property code_with_constants

Возвращает кортеж:

[0] красиво отформатированное представление (в виде корректного синтаксиса Python) внутреннего графа для метода forward. См. code. [1] ConstMap, следуя формату CONSTANT.cN вывода в [0]. Индексы в выводе [0] являются ключами к значениям соответствующих констант.

Подробнее см. Просмотр кода.

cpu()

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

Примечание

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

Возвращает:

self

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

Модуль

cuda(device=None)

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

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

Примечание

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

Параметры:

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

Возвращает:

self

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

Модуль

double()

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

Примечание

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

Возвращает:

self

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

Модуль

eval()

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

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

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

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

Возвращает:

self

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

Модуль

extra_repr()

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

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

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

str

float()

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

Примечание

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

Возвращает:

self

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

Модуль

get_buffer(target)

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

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

Параметры:

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

Возвращает:

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

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

torch.Tensor

Исключения:

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

get_extra_state()

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

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

Возвращает:

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

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

object

get_parameter(target)

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

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

Параметры:

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

Возвращает:

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

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

torch.nn.Parameter

Исключения:

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

get_submodule(target)

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

property graph

Возвращает строковое представление внутреннего графа для метода forward. См. Интерпретация графиков для получения подробностей.

half()

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

Примечание

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

Возвращает:

self

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

Модуль

property inlined_graph

Возвращает строковое представление внутреннего графа для метода forward. Этот граф будет предварительно обработан, чтобы встроить все вызовы функций и методов. См. Интерпретация графиков для получения подробностей.

ipu(device=None)

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

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

Примечание

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

Параметры:

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

Возвращает:

self

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

Модуль

load_state_dict(state_dict, strict=True)

Копирует параметры и буферы из 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()

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

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

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

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

Итератор[Модуль]

Примечание

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

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

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

(строка, тензор) – кортеж, содержащий имя и буфер

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

Итератор[Кортеж[строка, Тензор]]

Пример:

>>> for name, buf in self.named_buffers():
>>>    if name in ['running_var']:
>>>        print(buf.size())
named_children()

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

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

(строка, модуль) – кортеж, содержащий имя и дочерний модуль

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

Итератор[Кортеж[строка, Модуль]]

Пример:

>>> for name, module in model.named_children():
>>>     if name in ['conv4', 'conv5']:
>>>         print(module)
named_modules(memo=None, prefix='', remove_duplicate=True)

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

Параметры:
  • memo (Необязательно[Множество[Модуль]]) – запоминающее устройство для хранения набора уже добавленных модулей в результат
  • prefix (строка) – префикс, который будет добавлен к имени модуля
  • remove_duplicate (bool) – удалять ли дублированные экземпляры модулей в результате
Возвращаемые значения:

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

Примечание

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

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

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

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

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

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

Пример:

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

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

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

Параметры:

recurse (bool) – если 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)

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

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

Возвращает:

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

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

torch.utils.hooks.RemovableHandle

register_buffer(name, tensor, persistent=True)

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

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

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

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

Пример:

>>> self.register_buffer('running_mean', torch.zeros(num_features))
register_forward_hook(hook)

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

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

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

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

Возвращает:

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

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

torch.utils.hooks.RemovableHandle

register_forward_pre_hook(hook)

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

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

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

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

Возвращает:

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

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

torch.utils.hooks.RemovableHandle

register_full_backward_hook(hook)

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

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

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

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

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

Возвращает:

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

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

torch.utils.hooks.RemovableHandle

register_load_state_dict_post_hook(hook)

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

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

hook(module, incompatible_keys) -> None

Аргумент module — текущий модуль, для которого зарегистрирован этот хук, а аргумент incompatible_keys — NamedTuple, состоящий из атрибутов missing_keys и unexpected_keys. missing_keys — list str, содержащий пропущенные ключи, а unexpected_keys — list str, содержащий нежелательные ключи.

Заданные incompatible_keys могут быть изменены in-place, если это необходимо.

Обратите внимание, что проверки, выполняемые при вызове load_state_dict() с strict=True, зависят от изменений, которые хук вносит в missing_keys или unexpected_keys, как ожидается. Добавление в любой из наборов ключей приведёт к сгенерированной ошибке при strict=True, а очистка как пропущенных, так и нежелательных ключей позволит избежать ошибки.

Возвращает:

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

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

torch.utils.hooks.RemovableHandle

register_module(name, module)

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

register_parameter(name, param)

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

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

Параметры:
  • name (str) – имя параметра. К параметру можно получить доступ из этого модуля с помощью данного имени
  • param (Parameter или None) – параметр, который нужно добавить к модулю. Если None, операции, выполняемые над параметрами, такие как cuda, игнорируются. Если None, параметр не включён в state_dict модуля.
requires_grad_(requires_grad=True)

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

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

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

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

Параметры:

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

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

self

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

Модуль

save(f, _extra_files={})

См. torch.jit.save для подробностей.

set_extra_state(state)

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

Параметры:

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

share_memory()

См. torch.Tensor.share_memory_()

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

T

state_dict(*args, destination=None, prefix='', keep_vars=False)

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

Включаются как параметры, так и постоянные буферы (например, средние значения). Ключи соответствуют именам параметров и буферов. Параметры и буферы, установленные в None не включаются.

Примечание

Возвращаемый объект является поверхностной копией. Он содержит ссылки на параметры и буферы модуля.

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

В настоящее время state_dict() также принимает позиционные аргументы для destination, prefix и keep_vars в порядке следования. Однако это устаревает, и в будущих выпусках будут принудительно использоваться ключевые аргументы.

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

Пожалуйста, избегайте использования аргумента destination, так как он не предназначен для конечных пользователей.

Параметры:
  • destination (dict, необязательно) – Если указано, состояние модуля будет обновлено в словаре, и возвращается тот же объект. В противном случае создается и возвращается OrderedDict. По умолчанию: None.
  • prefix (str, необязательно) – префикс, добавляемый к именам параметров и буферов для составления ключей в state_dict. По умолчанию: ''.
  • keep_vars (bool, необязательно) – по умолчанию Tensor возвращаемые в state_dict открепляются от autograd. Если установлено в True, открепление не будет выполнено. По умолчанию: False.
Возвращаемое значение:

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

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

dict

Пример:

>>> module.state_dict().keys()
['bias', 'weight']
to(*args, **kwargs)

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

Его можно вызвать как

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

Его сигнатура похожа на torch.Tensor.to(), но принимает только вещественные или комплексные dtype. Кроме того, этот метод преобразует только вещественные или комплексные параметры и буферы в dtype (если указано). Целочисленные параметры и буферы будут перемещены device, если это указано, но с неизменными типами данных. Когда non_blocking установлен, он пытается асинхронно преобразовать/переместить с точки зрения хоста, если это возможно, например, перемещение тензоров CPU с закреплённой памятью на устройства CUDA.

Примеры см. ниже.

Примечание

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

Параметры:
  • device (torch.device) – желаемое устройство параметров и буферов в этом модуле
  • dtype (torch.dtype) – желаемый вещественный или комплексный тип данных параметров и буферов в этом модуле
  • tensor (torch.Tensor) – Тензор, тип данных и устройство которого являются желаемым типом данных и устройством для всех параметров и буферов в этом модуле
  • memory_format (torch.memory_format) – желаемый формат памяти для 4D параметров и буферов в этом модуле (только ключевой аргумент)
Возвращаемое значение:

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)

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

Параметры:

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

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

self

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

Модуль

END_OF_DOCUMENT_MARKER
train(mode=True)

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

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

Параметры:

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

Возвращает:

self

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

Модуль

type(dst_type)

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

Примечание

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

Параметры:

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

Возвращает:

self

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

Модуль

xpu(device=None)

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

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

Примечание

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

Параметры:

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

Возвращает:

self

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

Модуль

zero_grad(set_to_none=False)

Устанавливает градиенты всех параметров модели в ноль. См. аналогичную функцию в 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.jit.ScriptModule.html

Spec-Zone.ru

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