МодульScriptModule
-
class torch.jit.ScriptModule[исходный код] -
Обёртка вокруг C++
torch::jit::Module.ScriptModuleсодержат методы, атрибуты, параметры и константы. К ним можно обратиться так же, как и к обычнымnn.Module.-
add_module(name, 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) -
Перемещает все параметры и буферы модели на графический процессор.
Это также создаёт новые объекты для связанных параметров и буферов. Поэтому его следует вызывать до создания оптимизатора, если модуль будет жить на графическом процессоре во время оптимизации.
Примечание
Этот метод изменяет модуль на месте.
-
double() -
Преобразует все параметры и буферы с плавающей запятой в тип
double.Примечание
Этот метод изменяет модуль на месте.
- Возвращает:
-
self
- Тип возвращаемого значения:
-
eval() -
Устанавливает модуль в режим оценки.
Это оказывает влияние только на определённые модули. См. документацию по конкретным модулям для получения подробной информации об их поведении в режиме обучения/оценки, если они затронуты, например,
Dropout,BatchNorm, и т. д.Эквивалентно
self.train(False).См. Локальное отключение вычисления градиента для сравнения между
.eval()и несколькими аналогичными механизмами, которые могут быть с ним перепутаны.- Возвращает:
-
self
- Тип возвращаемого значения:
-
extra_repr() -
Устанавливает дополнительное представление модуля
Для вывода настраиваемой дополнительной информации, вы должны переопределить этот метод в собственных модулях. Разрешены как однострочные, так и многострочные строки.
- Тип возвращаемого значения:
-
float() -
Преобразует все параметры и буферы с плавающей запятой в тип
float.Примечание
Этот метод изменяет модуль на месте.
- Возвращает:
-
self
- Тип возвращаемого значения:
-
-
get_buffer(target) -
Возвращает буфер, заданный
target, если он существует, в противном случае генерирует ошибку.См. строку документации для
get_submoduleдля более подробного объяснения функциональности этого метода, а также того, как правильно указатьtarget.- Параметры:
-
target (str) – Полное квалифицированное строковое имя буфера для поиска. (См.
get_submoduleдля указания полного квалифицированного имени.) - Возвращает:
-
Буфер, на который ссылается
target - Тип возвращаемого значения:
- Исключения:
-
AttributeError – Если целевая строка ссылается на недопустимый путь или резолвится в элемент, который не является буфером
-
get_extra_state() -
Возвращает любое дополнительное состояние, которое следует включить в state_dict модуля. Реализуйте этот метод и соответствующий
set_extra_state()для вашего модуля, если вам нужно сохранить дополнительное состояние. Эта функция вызывается при построенииstate_dict()модуля.Обратите внимание, что дополнительное состояние должно быть сериализуемым с помощью pickle, чтобы гарантировать корректную сериализацию state_dict. Мы гарантируем обратную совместимость только для сериализации тензоров; другие объекты могут нарушить обратную совместимость, если их сериализованный вид с помощью pickle изменится.
- Возвращает:
-
Любое дополнительное состояние для хранения в state_dict модуля
- Тип возвращаемого значения:
-
get_parameter(target) -
Возвращает параметр, заданный
target, если он существует, в противном случае генерирует ошибку.См. строку документации для
get_submoduleдля более подробного объяснения функциональности этого метода, а также того, как правильно указатьtarget.- Параметры:
-
target (str) – Полное квалифицированное строковое имя параметра для поиска. (См.
get_submoduleдля указания полного квалифицированного имени.) - Возвращает:
-
Параметр, на который ссылается
target - Тип возвращаемого значения:
-
torch.nn.Parameter
- Исключения:
-
AttributeError – Если целевая строка ссылается на недопустимый путь или резолвится в элемент, который не является
nn.Parameter
-
get_submodule(target) -
Возвращает подмодуль, заданный
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
-
property graph -
Возвращает строковое представление внутреннего графа для метода
forward. См. Интерпретация графиков для получения подробностей.
-
half() -
Преобразует все параметры и буферы с плавающей запятой в тип данных
half.Примечание
Этот метод изменяет модуль на месте.
- Возвращает:
-
self
- Тип возвращаемого значения:
-
property inlined_graph -
Возвращает строковое представление внутреннего графа для метода
forward. Этот граф будет предварительно обработан, чтобы встроить все вызовы функций и методов. См. Интерпретация графиков для получения подробностей.
-
ipu(device=None) -
Перемещает все параметры и буферы модели на IPU.
Это также создаёт разные объекты для ассоциированных параметров и буферов. Поэтому его следует вызывать перед построением оптимизатора, если модуль будет жить на IPU во время оптимизации.
Примечание
Этот метод изменяет модуль на месте.
-
-
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) -
Возвращает итератор по буферам модуля, возвращая имя буфера и сам буфер.
- Параметры:
- Возвращаемые значения:
-
(строка, тензор) – кортеж, содержащий имя и буфер
- Тип возвращаемого значения:
Пример:
>>> 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) -
Возвращает итератор по параметрам модуля, возвращая имя параметра и сам параметр.
- Параметры:
- Возвращаемые значения:
-
(строка, параметр) – кортеж, содержащий имя и параметр
- Тип возвращаемого значения:
Пример:
>>> 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_meanBatchNorm не является параметром, но является частью состояния модуля. Буферы по умолчанию постоянны и будут сохранены вместе с параметрами. Это поведение можно изменить, установив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—liststr, содержащий пропущенные ключи, аunexpected_keys—liststr, содержащий нежелательные ключи.Заданные 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_()и несколькими похожими механизмами, с которыми его можно перепутать.
-
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
-
См.
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.
-
destination (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 параметров и буферов в этом модуле (только ключевой аргумент)
-
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) -
Перемещает параметры и буферы на указанное устройство без копирования хранилища.
- Параметры:
-
device (
torch.device) – Желаемое устройство параметров и буферов в этом модуле. - Возвращаемое значение:
-
self
- Тип возвращаемого значения:
-
-
train(mode=True) -
Устанавливает модуль в режим обучения.
Это оказывает влияние только на некоторые модули. Смотрите документацию конкретных модулей для получения подробностей о их поведении в режиме обучения/оценки, если они затронуты, например
Dropout,BatchNorm, и т.д.
-
type(dst_type) -
Преобразует все параметры и буферы в
dst_type.Примечание
Этот метод изменяет модуль на месте.
-
xpu(device=None) -
Перемещает все параметры и буферы модели на XPU.
Это также делает связанные параметры и буферы другими объектами. Поэтому его следует вызывать перед построением оптимизатора, если модуль будет существовать на XPU во время оптимизации.
Примечание
Этот метод изменяет модуль на месте.
-
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