拓冰建站拓冰建站
首页 / 资讯中心 / 正文

PyTorch torch.overrides 模块完全指南:深入理解与定制 __torch_function__ 协议

PyTorch torch.overrides 模块完全指南深入理解与定制torch_function协议【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorchtorch.overrides是 PyTorch 中为__torch_function__协议提供各类辅助函数的核心模块它定义了哪些 torch API 可以被 Tensor 子类或 Tensor-like 类型重载、如何在 Python 层面完成重载分发、以及如何绕过一层分发直接调用底层实现。本文以 torch.overrides.md 为主线结合 torch/overrides.py 的真实源码与 test/test_overrides.py 的测试用例系统讲解该模块每个公开函数的语义、底层实现与实战用法帮助你掌握在 PyTorch 中扩展自定义张量类型的标准姿势。一、模块定位Python 层的torch_function基础设施虽然__torch_function__协议的大部分处理逻辑位于 C 层但 PyTorch 的 API 中仍有相当一部分是用纯 Python 实现的因此需要 Python 层的配套处理。正如 torch/overrides.py 的文件头注释所说明的本文件是__torch_function__的 Python 实现。虽然大部分 torch API 与__torch_function__的处理发生在 C 层但部分 torch API 由 Python 编写因此也需要 Python 层的重载处理。面向开发者最主要的两个函数是handle_torch_function和has_torch_function使用示例见 torch/functional.py 与 test/test_overrides.py。该模块的设计深受 NumPy__array_function__协议NEP-0018启发。整个模块对外公开的 API 集合定义在 torch/overrides.py 的__all__中共 11 个符号__all__ [ get_ignored_functions, get_overridable_functions, get_testing_overrides, handle_torch_function, has_torch_function, resolve_name, is_tensor_like, is_tensor_method_or_property, wrap_torch_function, enable_reentrant_dispatch, redispatch_function, ]其中has_torch_function与handle_torch_function是最核心的开发者入口其余函数则服务于 API 内省、测试与诊断场景。关于__torch_function__协议的完整教程可参考 docs/source/notes/extending.md。二、API 内省三件套哪些函数可以被重载2.1get_ignored_functions()不可被重载的公开函数该函数返回一个set[Callable]包含所有公开暴露但在 torch API 中无法通过__torch_function__重载的函数。绝大多数情况下这些函数之所以不可重载是因为它们的参数中根本不包含 Tensor 或 Tensor-like 对象例如torch.is_tensor、torch.set_default_dtype、torch.get_num_threads、torch.device、torch.dtype等。实现位于 torch/overrides.py其关键特征有两个使用functools.cache缓存结果避免重复遍历构建集合的开销使用_disable_user_warnings装饰器临时屏蔽torch模块下形如xxx is deprecated, please use xxx的弃用警告保证函数无论何时调用都能返回一致的集合。原文档给出的两个判别示例 torch.Tensor.as_subclass in torch.overrides.get_ignored_functions() True torch.add in torch.overrides.get_ignored_functions() False从集合内容可以观察到几类典型成员见 torch/overrides.py类别典型成员不可重载原因全局配置类torch.set_default_tensor_type、torch.set_default_device、torch.manual_seed、torch.set_grad_enabled、torch.no_grad、torch.inference_mode参数与 Tensor 无关设备/类型查询类torch.has_cuda、torch.has_mps、torch.device、torch.dtype、torch.layout、torch.finfo操作对象不是 Tensor 实例工厂函数类torch.tensor、torch.as_tensor、torch.arange、torch.eye、torch.rand、torch.zeros、torch.empty、torch.linspace等由 C 层特判处理不参与 Python 层重载序列化/导入类torch.save、torch.load、torch.from_numpy、torch.frombuffer协议层功能Tensor 魔术方法Tensor.__init__、Tensor.__getattribute__、Tensor.__setattr__、Tensor.__torch_function__本身Python 对象模型机制Tensor 特殊方法Tensor.as_subclass、Tensor.cholesky、Tensor.eig、Tensor.lstsq、Tensor.qr、Tensor.solve及各类new_*工厂底层实现不经过 Python 分发autocast/确定性算法torch.set_autocast_enabled、torch.use_deterministic_algorithms、torch.set_float32_matmul_precision全局状态控制torch.nn.functional内部函数has_torch_function、handle_torch_function的 F 版本及_canonical_mask等内部基础设施2.2get_overridable_functions()列出所有可重载函数返回一个dict[Any, list[Callable]]以命名空间为键、该命名空间内所有可被__torch_function__重载的函数列表为值。其底层由_get_overridable_functions()torch/overrides.py实现遍历的命名空间包括tested_namespaces [ (torch, torch, torch.__all__), (torch.functional, torch.functional, torch.functional.__all__), (torch.nn.functional, torch.nn.functional, dir(torch.nn.functional)), (torch.nn.init, torch.nn.init, dir(torch.nn.init)), (torch.Tensor, torch.Tensor, dir(torch.Tensor)), (torch.linalg, torch.linalg, dir(torch.linalg)), (torch.fft, torch.fft, dir(torch.fft)), (torch.foreach, torch.foreach, torch.foreach.__all__), (torch.special, torch.special, dir(torch.special)), ]遍历时遵循严格的过滤规则跳过以__开头或_结尾的私有函数、跳过属性名不以小写字母开头的条目、跳过模块对象与__future__特性同时会与get_ignored_functions()交叉校验若某函数既在忽略集合中又存在显式重载则直接抛出AssertionError从机制上保证两份清单的一致性。函数使用functools.cache缓存并同样屏蔽弃用警告。2.3resolve_name(f)把函数解析为可读的字符串名给定一个传入__torch_function__的函数对象返回其人类可读的名称且该名称若被eval求值应能还原出原函数torch/overrides.py。对于torch._ops.OpOverload/OpOverloadPacket直接返回其str表示其余函数则通过_get_overridable_functions()[1]中维护的名称索引查询。名称索引形如torch.add、torch.Tensor.size属性访问器会额外登记__get__/__set__条目。三、get_testing_overrides()为测试生成全量哑重载get_testing_overrides()返回一个dict[Callable, Callable]将每一个可重载的 torch API 函数映射到一个签名相同、但无条件返回 -1 的 lambda。这些哑函数专门用于测试当一个自定义类型定义了__torch_function__时可用它验证该类型对 API 的覆盖是否完整。原文档示例torch/overrides.py import inspect my_add torch.overrides.get_testing_overrides()[torch.add] inspect.signature(my_add) Signature (input, other, outNone)由于 native kernel 无法被inspect直接解析签名见 Issue #28233 的注释这份字典由约 800 个手写 lambda 构成torch/overrides.py其中包含了大量 API 的完整默认参数例如torch.add: lambda input, other, outNone: -1, torch.addmm: lambda input, mat1, mat2, beta1, alpha1, outNone: -1, torch.allclose: lambda input, other, rtol1e-05, atol1e-08, equal_nanFalse: -1, torch.embedding_bag: lambda input, weight, offsets, max_normNone, norm_type2, scale_grad_by_freqFalse, modemean, sparseFalse, per_sample_weightsNone, padding_idxNone: -1, torch.nn.functional.scaled_dot_product_attention: lambda query, key, value, attn_maskNone, dropout_p0.0: -1,构建完成后还会做两类自动补全见 torch/overrides.py魔术方法与就地变体自动生成对每个基础函数按__name__、__name__ _、__ __name__ __、__i __name__ __、__r __name__ __五类命名在Tensor上查找对应属性bitwise_*系列还会额外生成__and__、__or__、__xor__等 dunder 形式。foreach 与分布式函数补充torch.foreach.*系列torch/overrides.py与torch.distributed的broadcast、all_reduce、all_gather、reduce_scatter等torch/overrides.py在自动生成循环之后单独追加以避免误生成无关的 Tensor 方法例如dist.reduce若走自动生成会衍生出Tensor.__reduce__。四、核心分发入口has_torch_function 与 handle_torch_function4.1has_torch_function(relevant_args)判断是否需要进行分发该函数实际是对 C 实现_has_torch_function的包装torch/overrides.py文档字符串明确了两条使用准则用途作为调用handle_torch_function之前的守卫判断检查可迭代对象relevant_args的元素中是否存在__torch_function__实现或当前是否启用了__torch_function__mode禁止用途不要用它判断某个对象是否是 Tensor-like——那是is_tensor_like的职责。特别地精确的torch.Tensor与torch.nn.Parameter实例被视为不可分发对象。此外模块还提供了两个性能优化变体has_torch_function_unary(t)针对单输入的特化省去元组打包/解包开销torch/overrides.pyhas_torch_function_variadic(a, b, ...)基于 Python 3.7 的METH_FASTCALL协议跳过元组创建直接传参torch/overrides.py。对应的 C 声明可见 torch/_C/init.pyi.in。测试 test/test_overrides.py 中专门覆盖了test_has_torch_function_non_sequence等边界场景。4.2handle_torch_function(public_api, relevant_args, *args, **kwargs)执行分发这是 Python 层分发的核心函数torch/overrides.py其 C 对等物是torch::autograd::handle_torch_function。调用约定为把最初被调用的公开 API 函数作为public_api传入relevant_args是需要检查的 Tensor-like 参数集合args/kwargs是原始调用参数。典型的手写分发函数模式torch/overrides.pydef func(a): if has_torch_function_unary(a): return handle_torch_function(func, (a,), a) return a 0其分发流程分为两步第一步收集重载参数。_get_overloaded_argstorch/overrides.py遍历relevant_args只收集类型唯一且定义了__torch_function__且未被_disabled_torch_function_impl禁用的参数并按照 NEP-0018 描述的优先级算法排序子类优先于父类同层级则按参数从左到右的顺序。其复杂度为 O(参数个数 × 唯一类型数)。第二步依次调用重载实现。对每个重载参数按优先级调用其__torch_function__(public_api, types, args, kwargs)若某实现返回非NotImplemented的值立即返回该值若返回NotImplemented则继续尝试下一个参数的重载全部返回NotImplemented时抛出TypeError错误信息形如no implementation found for torch.add on types that implement __torch_function__: [ScalarTensor]若启用了 mode 还会附加当前 mode 信息torch/overrides.py。在 Python 3.14 之前将 __torch_function__ 定义为普通实例方法会触发 DeprecationWarning官方要求将其定义为 classmethod见 [torch/overrides.py](https://link.gitcode.com/i/994ef44d4852a3f55b89bb80a2a0b321#L1876-L1890)。4.3 TorchFunctionMode无需子类的全局重载除参数重载外handle_torch_function还会在调用参数重载之前检查_is_torch_function_mode_enabled()torch/overrides.py若启用了 mode则临时弹出栈顶 mode 并调用其__torch_function__若结果非NotImplemented则直接返回。TorchFunctionMode类torch/overrides.py适用于三类场景重载工厂函数等本就不接收 Tensor 参数的函数这些无法通过 Tensor 子类重载希望在不包装输入的前提下重载所有函数例如只想记录中间计算需要显式控制多个 Tensor 子类的执行顺序而非隐式依赖NotImplemented的返回值。独立子类之间具有组合性with MyMode():会 push 到模式栈上在__torch_function__实现内继续调用 torch API 默认会转发给栈中的下一个 mode。五、类型判断与元信息辅助函数5.1is_tensor_like(inp)判断是否为 Tensor-like实现非常简洁torch/overrides.pyreturn type(inp) is torch.Tensor or hasattr(inp, __torch_function__)即精确的Tensor实例或类型上带有__torch_function__属性的对象都算 Tensor-like。原文档给出的判别矩阵 class SubTensor(torch.Tensor): ... is_tensor_like(SubTensor([0])) # Tensor 子类通常是 Tensor-like True is_tensor_like(6) # 内建类型不是 False is_tensor_like(None) # None 不是 False class NotATensor: ... is_tensor_like(NotATensor()) # 普通用户类型不是 False class TensorLike: ... classmethod ... def __torch_function__(cls, func, types, args, kwargs): ... return -1 is_tensor_like(TensorLike()) # 实现 __torch_function__ 后变成 Tensor-like True5.2is_tensor_method_or_property(func)识别 Tensor 方法/属性处理器__torch_function__收到的func不仅可能是模块级函数也可能是torch.Tensor的方法或属性属性传入的是其__get__。该函数用于判断这一点torch/overrides.py实现方式是在get_overridable_functions()[torch.Tensor]集合中查找。之所以需要专门的判断是因为方法/属性有时没有__module__槽位且它们要求第一个参数必须是torch.Tensor实例。 is_tensor_method_or_property(torch.Tensor.add) True is_tensor_method_or_property(torch.add) False六、函数装饰器与重分发机制6.1wrap_torch_function(dispatcher)把任意函数变成可分发函数这是一个装饰器工厂torch/overrides.py。dispatcher是与被装饰函数同签名的可调用对象负责从参数中抽取 Tensor-like 集合。其内部包装逻辑为def wrapped(*args, **kwargs): relevant_args dispatcher(*args, **kwargs) if has_torch_function(relevant_args): return handle_torch_function(wrapped, relevant_args, *args, **kwargs) return func(*args, **kwargs)原文档示例def dispatcher(a): # 必须与 func 签名一致 return (a,) torch.overrides.wrap_torch_function(dispatcher) def func(a): # 使 func 可被 __torch_function__ 分发 return a 0需要注意的是文档明确警告该装饰器可能降低性能通常把代码写成一系列本身支持__torch_function__的函数就足够了只有包装底层库且同时要求支持 Tensor-like 的罕见场景才需要它。test/test_overrides.py中的test_wrap_torch_functiontest/test_overrides.py提供了验证用例。6.2redispatch_function(func, types, args, kwargs)跳过一层分发该函数用于跳过一层__torch_function__分发直接调用函数实现torch/overrides.py底层委托给 C 的_skip_one_hop_torch_function见 torch/_C/init.pyi.in。它主要服务于希望调用函数实现、同时仍拦截该函数内部 PyTorch 操作的 Tensor 子类。原文档给出了完整的LoggingTensor示例子类在__torch_function__中打印调用名后通过redispatch_function继续执行可以观察到外层scaled_mul与内部mul都被记录而 1这一加法不会触发日志——因为redispatch_function返回的是普通torch.Tensor。若改用TorchFunctionMode并在redispatch_function外层配合with self:mode 会跨内部所有操作保持激活此时add也能被记录。两者行为差异恰好说明了参数分发与mode 分发的本质区别。七、相关基础设施enable_reentrant_dispatch 与分发状态enable_reentrant_dispatch()是torch._C._RestorePythonTLSSnapshot的上下文管理器包装torch/overrides.py。由于torch._C._RestorePythonTLSSnapshot在模块导入初期尚不可用且直接赋值会改变__module__使其看起来像私有 API因此以公开函数的形式提供。模块还围绕 mode 栈提供了底层辅助_get_current_function_mode()、_get_current_function_mode_stack()、_push_mode()、_pop_mode()与_pop_mode_temporarily()torch/overrides.py分别对应 C 层的_len_torch_function_stack、_get_function_stack_at、_push_on_torch_function_stack、_pop_torch_function_stack。八、实战从零实现一个支持torch_function的类型结合 docs/source/notes/extending.md 的经典ScalarTensor示例可以看到torch.overrides模块在真实项目中的完整用法。第一步定义带__torch_function__的类与分发表import functools import torch HANDLED_FUNCTIONS {} class ScalarTensor(object): def __init__(self, N, value): self._N N self._value value def __repr__(self): return ScalarTensor(N{}, value{}).format(self._N, self._value) def tensor(self): return self._value * torch.eye(self._N) classmethod def __torch_function__(cls, func, types, args(), kwargsNone): if kwargs is None: kwargs {} if func not in HANDLED_FUNCTIONS or not all( issubclass(t, (torch.Tensor, ScalarTensor)) for t in types ): return NotImplemented return HANDLED_FUNCTIONSfunc第二步用装饰器注册重载def implements(torch_function): Register a torch function override for ScalarTensor def decorator(func): functools.update_wrapper(func, torch_function) HANDLED_FUNCTIONS[torch_function] func return func return decorator implements(torch.mean) def mean(input): return float(input._value) / input._N implements(torch.add) def add(input, other): try: if input._N other._N: return ScalarTensor(input._N, input._value other._value) else: raise ValueError(Shape mismatch!) except AttributeError: return torch.add(ensure_tensor(input), ensure_tensor(other))第三步验证分发行为 d ScalarTensor(5, 2) torch.mean(d) 0.4 torch.add(d, d) ScalarTensor(N2, value4) torch.mul(d, 3) # 未注册的重载 → NotImplemented → TypeError TypeError: no implementation found for torch.mul on types that implement __torch_function__: [ScalarTensor]需要注意__torch_function__分发机制不校验重载函数与原始函数的签名一致性。例如torch.add支持alpha关键字参数若重载的add未定义该参数torch.add(s, s, alpha2)会直接抛出TypeError: add() got an unexpected keyword argument alpha。因此为保证与Tensor的完全兼容重载实现应精确模拟被重载函数的 API。此外协议设计目标是全量覆盖 API——部分覆盖可能导致某些函数抛出TypeError甚至无限递归对Tensor子类而言torch.add、torch.Tensor.__add__、torch.Tensor.add三处都需要覆盖且实现内部应调用super().__torch_function__(...)而非直接调用func。九、性能基线torch_function的开销度量由于handle_torch_function位于每个可重载调用的热路径上仓库在 benchmarks/overrides_benchmark 提供了专门的微基准套件其使用说明见 benchmarks/overrides_benchmark/README.md# 需先安装 py-spy pip install py-spy # 在 benchmarks/overrides_benchmark 目录下运行 python bench.py # 基准全部场景 py-spy record -o tensor.svg --native -- python pyspybench.py Tensor py-spy record -o overridden.svg --native -- python pyspybench.py WithTorchFunction基准预期结果README 原文要点对普通torch.Tensor输入执行 torch 函数的开销约为 2 μs__torch_function__应对torch.Tensor输入零开销对torch.Tensor子类有少量开销对定义__torch_function__的 Tensor-like 有几微秒开销分发机制的小幅改动约 100 ns 量级难以从噪声中分辨但对性能至关重要。因此 torch/overrides.py 明确要求任何可能影响__torch_function__开销的改动都必须上报该目录下的基准结果。十、测试与验证test/test_overrides.py模块行为由 test/test_overrides.py 系统性验证测试覆盖了语义正确性test_dtype_override、test_mean_semantics、test_mm_semantics、test_precedence_semantics多类型重载优先级、test_user_implementation_raises参数探测test_has_torch_function_non_sequence、test_torch_function_in_lists、test_torch_function_in_float_lists、test_torch_function_in_scalar_lists、test_torch_function_precedence_in_lists、test_torch_function_mixed_lists、test_torch_function_empty_lists子类行为test_tensor_subclass_propagation、test_base、test_grad、test_getitem_subclass基础设施test_wrap_torch_function、test_wrapper、test_broadcast_all。结语torch.overrides是 PyTorch 扩展机制的关键枢纽get_ignored_functions与get_overridable_functions划定了可重载边界has_torch_function/handle_torch_function构成了分发主路径redispatch_function与TorchFunctionMode提供了细粒度控制get_testing_overrides则为自定义类型的 API 覆盖测试提供了免费的全量哑重载。理解这一模块是写出行为正确、性能可控的 Tensor 子类与 Tensor-like 类型的前提。若需进一步深入建议依次阅读 docs/source/notes/extending.md协议完整教程、torch/functional.pyhandle_torch_function的生产级用法与 test/test_overrides.py行为契约。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门