Spec-Zone.ru › PyTorch 2.14

torch.overrides

Создано: 30 нояб. 2020 | Последнее обновление: 10 мая 2026

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

Функции

torch.overrides.get_ignored_functions() [исходный код]

Возвращает общедоступные функции, которые нельзя переопределить с помощью __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() [исходный код]

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

Возвращает:

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

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

Dict[Any, List[Callable]]

torch.overrides.resolve_name(f) [исходный код]

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

Параметры:

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

Возвращает:

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

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

str

torch.overrides.get_testing_overrides() [исходный код]

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

Возвращает:

Словарь, сопоставляющий функции 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) [исходный код]

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

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

Параметры:
  • 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) [исходный код]

Возвращает 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) [исходный код]

Возвращает 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) [исходный код]

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

Параметры:

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

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

Callable[[Callable[[~_P], _R]], Callable[[~_P], _R]]

Примечание

Этот декоратор может снизить производительность кода. Как правило, достаточно представить код в виде последовательности функций, которые сами поддерживают __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
torch.overrides.redispatch_function(func, types, args, kwargs) [исходный код]

Пропускает один уровень диспетчеризации __torch_function__ и вызывает функцию.

Это в первую очередь полезно для подклассов Tensor, которым нужно вызвать реализацию функции, продолжая перехватывать операции PyTorch внутри неё.

Пример с подклассом Tensor. Перехватываются только операции, среди входных данных которых есть LoggingTensor; когда redispatch_function возвращает обычный torch.Tensor, последующие операции (здесь + 1) уже не регистрируются.

>>> from torch.overrides import has_torch_function, handle_torch_function
>>> class LoggingTensor(torch.Tensor):
...     depth = 0
...
...     @classmethod
...     def __torch_function__(cls, func, types, args, kwargs=None):
...         print(f"{'  ' * cls.depth}Calling {func.__name__}")
...         cls.depth += 1
...         r = torch.overrides.redispatch_function(func, types, args, kwargs)
...         cls.depth -= 1
...         return r
>>> def scaled_mul(a, b):
...     if has_torch_function((a, b)):
...         return handle_torch_function(scaled_mul, (a, b), a, b)
...     return a * b + 1
>>> x = LoggingTensor(torch.tensor([3.0]))
>>> y = LoggingTensor(torch.tensor([4.0]))
>>> result = scaled_mul(x, y)
Calling scaled_mul
  Calling mul
>>> result
tensor([13.])

Обратите внимание: регистрируется только mul, но не add: redispatch_function возвращает обычный torch.Tensor, поэтому + 1 внутри scaled_mul больше не видит входной объект LoggingTensor, и __torch_function__ не вызывается.

При использовании TorchFunctionMode режим остаётся активным для всех внутренних операций, поэтому теперь видна и + 1. Используйте with self: после redispatch_function, чтобы повторно включить режим для этих внутренних вызовов.

>>> from torch.overrides import TorchFunctionMode
>>> class LoggingMode(TorchFunctionMode):
...     def __init__(self):
...         self.depth = 0
...
...     def __torch_function__(self, func, types, args, kwargs=None):
...         print(f"{'  ' * self.depth}Calling {func.__name__}")
...         self.depth += 1
...         with self:
...             r = torch.overrides.redispatch_function(
...                 func, types, args, kwargs
...             )
...         self.depth -= 1
...         return r
>>> a = torch.tensor([3.0])
>>> b = torch.tensor([4.0])
>>> with LoggingMode():
...     result = scaled_mul(a, b)
Calling scaled_mul
  Calling mul
  Calling add
>>> result
tensor([13.])

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

Spec-Zone.ru

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