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) – Функция, для которой требуется получить имя.
- Возвращает
-
Имя функции; если оно будет вычислено, оно должно вернуть исходную функцию.
- Тип возвращаемого значения
-
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__.:type relevant_args: iterable- Возвращает
-
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/2.1/torch.overrides.html