ScriptModule
-
class torch.jit.ScriptModule[source] -
Оборачивающий класс для 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] являются ключами к значениям констант.Подробности см. в разделе Просмотр кода.
-
compile(*args, **kwargs) -
Компилирует метод forward этого модуля с использованием
torch.compile().Метод
__call__этого модуля компилируется, и все аргументы передаются как есть вtorch.compile().Подробности об аргументах для этой функции см. в
torch.compile().
-
cpu() -
Перемещает все параметры и буферы модели на CPU.
Примечание
Этот метод изменяет модуль на месте.
- Возвращает
-
self
- Тип возвращаемого значения
-
cuda(device=None) -
Перемещает все параметры и буферы модели на GPU.
Это также создает разные объекты для связанных параметров и буферов. Поэтому его следует вызывать перед созданием оптимизатора, если модуль будет жить на GPU во время оптимизации.
Примечание
Этот метод изменяет модуль на месте.
-
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()модуля.Обратите внимание, что дополнительные данные должны быть сериализуемыми, чтобы гарантировать правильную сериализацию 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, 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() -
Возвращает итератор по всем модулям в сети.
Примечание
Дублируемые модули возвращаются только один раз. В следующем примере
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) — Кортеж, содержащий имя и буфер
- Тип возвращаемого значения
Пример:
>>> for name, buf in self.named_buffers(): >>> if name in ['running_var']: >>> print(buf.size())
-
named_children() -
Возвращает итератор по непосредственным дочерним модулям, возвращая имя модуля и сам модуль.
- Выход
-
(str, Модуль) — Кортеж, содержащий имя и дочерний модуль
- Тип возвращаемого значения
Пример:
>>> for name, module in model.named_children(): >>> if name in ['conv4', 'conv5']: >>> print(module)
-
named_modules(memo=None, prefix='', remove_duplicate=True) -
Возвращает итератор по всем модулям в сети, возвращая имя модуля и сам модуль.
- Параметры
- Выход
-
(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, Параметр) — Кортеж, содержащий имя и параметр
- Тип возвращаемого значения
Пример:
>>> for name, param in self.named_parameters(): >>> if name in ['bias']: >>> print(param.size())
-
-
parameters(recurse=True) -
Возвращает итератор по параметрам модуля.
Обычно передаётся в оптимизатор.
- Параметры
-
recurse (bool) – если True, то возвращает параметры этого модуля и всех его подмодулей. В противном случае, возвращает только параметры, являющиеся прямыми членами этого модуля.
- Возвращает
-
Parameter – параметр модуля
- Тип возвращаемого значения
Пример:
>>> for param in model.parameters(): >>> print(type(param), param.size()) <class 'torch.Tensor'> (20L,) <class 'torch.Tensor'> (20L, 1L, 5L, 5L)
-
register_backward_hook(hook) -
Регистрирует обратный хук для модуля.
Эта функция устарела в пользу
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— список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) -
Псевдоним для
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_()и несколькими похожими механизмами, которые могут быть с ним перепутаны.
-
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
-
-
См.
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.
-
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, recurse=True) -
Перемещает параметры и буферы на указанное устройство без копирования хранилища.
- Параметры
-
-
device (
torch.device) – Желаемое устройство параметров и буферов в этом модуле. - recurse (bool) – Перемещать ли параметры и буферы дочерних модулей рекурсивно на указанное устройство.
-
device (
- Возвращаемое значение
-
self
- Тип возвращаемого значения
-
train(mode=True) -
Переводит модуль в режим обучения.
Это оказывает влияние только на определённые модули. См. документацию конкретных модулей для подробностей о их поведении в режимах обучения/оценки, если они затрагиваются, например,
Dropout,BatchNorm, и т.д.
-
type(dst_type) -
Преобразует все параметры и буферы в
dst_type.Примечание
Этот метод изменяет модуль на месте.
-
xpu(device=None) -
Перемещает все параметры и буферы модели на XPU.
Также создаёт новые объекты для сопутствующих параметров и буферов. Поэтому его нужно вызывать перед построением оптимизатора, если модуль будет существовать на XPU во время оптимизации.
Примечание
Этот метод изменяет модуль на месте.
-
-
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