Spec-Zone.ru › PyTorch 2

torch.overrides

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

Функции

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__

Параметры

f (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__.:type relevant_args: iterable

Возвращает

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/2.1/torch.overrides.html

Spec-Zone.ru

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