Spec-Zone.ru › PyTorch 1

torch.overrides

Этот модуль предоставляет различные вспомогательные функции для протокола __torch_function__. Подробнее о протоколе __torch_function__ см. в Расширение torch.

Функции

torch.overrides.get_ignored_functions() [source]

Возвращает публичные функции, которые нельзя переопределить с помощью __torch_function__.

Возвращает:

Кортеж функций, которые общедоступны в API torch, но не могут быть переопределены с помощью __torch_function__. В основном это связано с тем, что ни один из аргументов этих функций не является тензором или тензорно-подобным объектом.

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

Set[Callable]

Примеры

>>> torch.Tensor.as_subclass in torch.overrides.get_ignored_functions()
True
>>> torch.add in torch.overrides.get_ignored_functions()
False
torch.overrides.get_overridable_functions() [source]

Список функций, которые можно переопределить с помощью __torch_function__

Возвращает:

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

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

Dict[Any, List[Callable]]

torch.overrides.resolve_name(f) [source]

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

Параметры:

callable (Callable) – Функция, для которой нужно получить имя.

Возвращает:

Имя функции; при вычислении оно должно вернуть исходную функцию.

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

str

torch.overrides.get_testing_overrides() [source]

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

Возвращает:

Словарь, который сопоставляет переопределяемые функции в API PyTorch с лямбда-функциями, имеющими ту же сигнатуру, что и реальная функция, и безусловно возвращающими -1. Эти лямбда-функции полезны для тестирования охвата API для типа, который определяет __torch_function__.

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

Dict[Callable, Callable]

Примеры

>>> import inspect
>>> my_add = torch.overrides.get_testing_overrides()[torch.add]
>>> inspect.signature(my_add)
<Signature (input, other, out=None)>
torch.overrides.handle_torch_function(public_api, relevant_args, *args, **kwargs) [source]

Реализует функцию с проверками на переопределения __torch_function__.

См. torch::autograd::handle_torch_function для эквивалента этой функции в реализации на C++.

Параметры:
  • public_api (function) – Функция, экспонируемая общедоступным API torch, которая изначально вызывалась как public_api(*args, **kwargs), на аргументах которой сейчас проверяются.
  • relevant_args (iterable) – Итерируемый объект аргументов для проверки методов __torch_function__.
  • args (tuple) – Произвольные позиционные аргументы, изначально переданные в public_api.
  • kwargs (tuple) – Произвольные ключевые аргументы, изначально переданные в public_api.
Возвращает:

Результат вызова implementation или метода __torch_function__ по мере необходимости.

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

object

:raises TypeError : если не найдено никакого имплементации.

Пример

>>> def func(a):
...     if has_torch_function_unary(a):
...         return handle_torch_function(func, (a,), a)
...     return a + 0
torch.overrides.has_torch_function()

Проверка наличия имплементаций __torch_function__ в элементах итерируемого объекта или если включён режим __torch_function__. Рассматривает точные Tensor и Parameter как неподдерживаемые. Используйте это для защиты вызова handle_torch_function(); не используйте для проверки того, является ли что-то тензорно-подобным, используйте is_tensor_like() вместо этого. :param relevant_args: Итерируемый объект аргументов для проверки методов __torch_function__.

Возвращает:

True, если у какого-либо элемента relevant_args есть реализация __torch_function__, иначе False.

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

bool

См. также

torch.is_tensor_like

Проверяет, является ли что-то тензорно-подобным, включая точный Tensor.

torch.overrides.is_tensor_like(inp) [source]

Возвращает True если переданный вход является тензорно-подобным.

В настоящее время это происходит всякий раз, когда есть атрибут __torch_function__ в типе ввода.

Примеры

Подкласс тензора обычно является тензорно-подобным.

>>> class SubTensor(torch.Tensor): ...
>>> is_tensor_like(SubTensor([0]))
True

Встроенные или пользовательские типы обычно не являются тензорно-подобными.

>>> is_tensor_like(6)
False
>>> is_tensor_like(None)
False
>>> class NotATensor: ...
>>> is_tensor_like(NotATensor())
False

Однако их можно сделать тензорно-подобными, реализовав __torch_function__.

>>> class TensorLike:
...     @classmethod
...     def __torch_function__(cls, func, types, args, kwargs):
...         return -1
>>> is_tensor_like(TensorLike())
True
torch.overrides.is_tensor_method_or_property(func) [source]

Возвращает True, если переданная функция является обработчиком метода или свойства, принадлежащего torch.Tensor, как передано в __torch_function__.

Примечание

Для свойств должен быть передан метод __get__.

Это может потребоваться, в частности, по следующим причинам:

  1. Методы/свойства иногда не содержат слота __module__.
  2. Они требуют, чтобы первый переданный аргумент был экземпляром torch.Tensor.

Примеры

>>> is_tensor_method_or_property(torch.Tensor.add)
True
>>> is_tensor_method_or_property(torch.add)
False
Тип возвращаемого значения:

bool

torch.overrides.wrap_torch_function(dispatcher) [source]

Обертывает заданную функцию с функциональностью, связанной с __torch_function__.

Параметры:

dispatcher (Callable) – вызываемый объект, который возвращает итерируемый объект тензорно-подобных объектов, переданных в функцию.

Примечание

Этот декоратор может снизить производительность вашего кода. Обычно достаточно выразить ваш код как серию функций, которые сами поддерживают __torch_function__. Если вы столкнулись с редким случаем, когда это не так, например, если вы обертываете библиотеку низкого уровня, и вам также нужно, чтобы она работала с тензорно-подобными объектами, то эта функция доступна.

Примеры

>>> def dispatcher(a): # Must have the same signature as func
...     return (a,)
>>> @torch.overrides.wrap_torch_function(dispatcher)
>>> def func(a): # This will make func dispatchable by __torch_function__
...     return a + 0

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/torch.overrides.html

Spec-Zone.ru

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