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) – Функция, для которой нужно получить имя.
- Возвращает:
-
Имя функции; при вычислении оно должно вернуть исходную функцию.
- Тип возвращаемого значения:
-
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.
-
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__.- Возвращает:
-
True, если у какого-либо элемента relevant_args есть реализация __torch_function__, иначе False.
- Тип возвращаемого значения:
См. также
-
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__.Это может потребоваться, в частности, по следующим причинам:
- Методы/свойства иногда не содержат слота
__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)[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