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) – Функция, имя которой нужно определить.
- Возвращает:
-
Имя функции; при вычислении оно должно возвращать исходную функцию.
- Тип возвращаемого значения:
-
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.
-
public_api (function) – Функция общедоступного API torch, изначально вызывавшаяся как
- Возвращает:
-
Результат вызова
implementationили метода__torch_function__, в зависимости от ситуации. - Тип возвращаемого значения:
: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.
- Тип возвращаемого значения:
См. также
-
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__.Это может быть необходимо, в частности, по следующим причинам:
- У методов и свойств иногда отсутствует слот
__module__. - Они требуют, чтобы первый переданный аргумент был экземпляром
torch.Tensor.
Примеры
>>> is_tensor_method_or_property(torch.Tensor.add) True >>> is_tensor_method_or_property(torch.add) False
- Тип возвращаемого значения:
- У методов и свойств иногда отсутствует слот
-
torch.overrides.wrap_torch_function(dispatcher)[исходный код] -
Оборачивает указанную функцию, добавляя функциональность, связанную с
__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
-
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