Spec-Zone.ru › PyTorch 2

ScriptModule

class torch.jit.ScriptModule [source]

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

add_module(name, module)

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

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

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

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

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

Примечание

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

Возвращает

self

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

Module

buffers(recurse=True)

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

Параметры

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()

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

Возвращает

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

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

Iterator[Module]

property code

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

property code_with_constants

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

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

Подробности см. в разделе Просмотр кода.

compile(*args, **kwargs)

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

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

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

cpu()

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

Примечание

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

Возвращает

self

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

Module

cuda(device=None)

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

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

Примечание

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

Параметры

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

Возвращает

self

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

Module

double()

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

Примечание

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

Возвращает

self

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

Module

eval()

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

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

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

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

Возвращает

self

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

Module

extra_repr()

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

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

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

str

float()

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

Примечание

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

Возвращает

self

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

Module

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

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

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

Module

property inlined_graph

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

ipu(device=None)

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

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

Примечание

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

Параметры

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

Возвращает

self

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

Module

load_state_dict(state_dict, strict=True, assign=False)

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

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

Выход

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

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

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, remove_duplicate=True)

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

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

(str, torch.Tensor) — Кортеж, содержащий имя и буфер

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

Iterator[Tuple[str, Тензор]]

Пример:

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

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

Выход

(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)

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

Параметры
  • 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, remove_duplicate=True)

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

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

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

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

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

Пример:

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

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

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

Параметры

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)

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

Эта функция устарела в пользу 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 or None) – регистрируемый буфер. Если None, операции, выполняемые над буферами, такие как cuda, игнорируются. Если None, буфер не включается в state_dict модуля.
  • persistent (bool) – является ли буфер частью state_dict этого модуля.

Пример:

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

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

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

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

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

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

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

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

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

torch.utils.hooks.RemovableHandle

register_forward_pre_hook(hook, *, prepend=False, with_kwargs=False)

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

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

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

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

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

hook(module, args, kwargs) -> None or a tuple of modified input and kwargs
Параметры
  • hook (Callable) – пользовательский хук для регистрации.
  • prepend (bool) – Если true, предоставленный 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) – Если true, hook будет получать ключевые аргументы, переданные функции forward. По умолчанию: False
Возвращает

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

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

torch.utils.hooks.RemovableHandle

register_full_backward_hook(hook, prepend=False)

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

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

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

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

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

Параметры
  • hook (Callable) – Пользовательский хук для регистрации.
  • prepend (bool) – Если True, предоставленный 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)

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

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

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

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

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

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

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

Параметры
  • hook (Callable) – Пользовательский хук для регистрации.
  • prepend (bool) – Если True, предоставленный 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)

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

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

hook(module, incompatible_keys) -> None

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

Указанный incompatible_keys можно изменять на месте, если необходимо.

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

Возвращает

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

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

torch.utils.hooks.RemovableHandle

register_module(name, module)

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

register_parameter(name, param)

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

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

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

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

requires_grad_(requires_grad=True)

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

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

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

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

Параметры

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

Возвращает

self

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

Module

save(f, _extra_files={})

См. torch.jit.save, который принимает файловый объект. Эта функция torch.save() преобразует объект в строку, рассматривая его как путь. НЕ путайте эти две функции, когда речь идёт о функциональности параметра ‘f’.

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 открепляются от автографа. Если установлено в 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, recurse=True)

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

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

self

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

Модуль

train(mode=True)

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

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

Параметры

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

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

self

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

Модуль

type(dst_type)

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

Примечание

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

Параметры

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

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

self

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

Модуль

xpu(device=None)

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

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

Примечание

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

Параметры

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

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

self

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

Модуль

zero_grad(set_to_none=True)

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

Spec-Zone.ru

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