PyTorch FakeTensor 深度解析:不执行计算即可获得张量元数据(shape/dtype/device)的模拟机制

发布时间:2026/9/10 12:13:41
PyTorch FakeTensor 深度解析:不执行计算即可获得张量元数据(shape/dtype/device)的模拟机制 PyTorch FakeTensor 深度解析不执行计算即可获得张量元数据shape/dtype/device的模拟机制【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorchFakeTensor 是 PyTorch 2.x 编译栈Dynamo / AOTAutograd / Inductor中用于无数据张量模拟的核心基础设施它能在不分配实际存储、不执行任何数值计算的前提下精确回答某个算子会输出什么 shape、dtype、device、以及是否存在别名关系等问题。本文将基于 docs/source/user_guide/torch_compiler/torch.compiler_fake_tensor.md 官方指南结合仓库源码torch/_subclasses/fake_tensor.py、torch/_guards.py、torch/fx/passes/fake_tensor_prop.py与测试用例test/test_fake_tensor.py完整讲解其动机、架构、API、实现原理与性能特征帮助你理解乃至亲手使用这套元数据模拟机制。为什么需要 FakeTensor动机在 Dynamo 符号求值symbolic evaluation和编译器优化 pass 中我们经常需要执行张量算子来获知输出的 shape、dtype、device——但又不希望真正跑这些算子。理由很直接速度大量真实计算会显著拖慢编译过程内存编译器在编译程序期间占用 GPU 显存是非常糟糕的尤其是你正在为 GPU 程序做分析时。官方指南举了一个典型场景Dynamo 跟踪用户 Tensor 代码时需要回答关于中间张量的查询例如用户对中间张量做条件判断。没有 FakeTensor这些查询就无法获得准确信息。另一个典型场景是把张量元数据存储在 FX IR 节点上即meta[val]。直接在节点上挂一个 fake tensor就能获得所需的全部元数据包括那些很容易被遗漏的细节——例如别名关系aliasing relationships。这是手写元数据传播逻辑很难做到的事情。FakeTensor 的定义与官方描述一致fake tensor 在所有方面都像一个真实张量唯一区别是它不真正持有数据。与相关概念的辨析为了准确理解 FakeTensor官方指南先厘清了三个相邻概念Meta Tensor同门但不足Meta tensor是devicemeta的张量它已经满足了 FakeTensor 的大部分需求不存数据、跑 meta kernel。但存在两个关键短板meta tensor 不建模设备stride 行为有时随设备不同而变化例如 CPU 与 CUDA 的某些布局差异FakeTensor 通过携带fake_device能获得更精确的信息生命周期作用域meta tensor 是全局存在的和普通 CPU/CUDA 张量一样独立存在而 fake tensor 被限定在某个FakeTensorMode的作用域内。Tensor SubclassFakeTensor 的实现载体tensor subclass允许你继承torch.Tensor并定制其行为。FakeTensor 正是以 tensor subclass 形式实现的——这意味着它几乎全部实现都写在 Python 里。这也解释了为什么 FakeTensor 的调试、扩展相对容易但也带来了后续会提到的性能开销。Dynamic Shapes与 ShapeEnv 的绑定动态形状dynamic shapes允许张量使用符号尺寸symbolic sizes而非具体尺寸并在算子中符号化地传播。动态形状的状态维护在ShapeEnv中而ShapeEnv 始终与某个FakeTensorMode关联——所以 FakeTensor 同时承担着管理符号尺寸的职责。每当用 PT2 编译一个子图时都会存在一个与该编译绑定的 tracing context其中包含除其他内容外一个FakeTensorMode以及可能一个ShapeEnv。整体架构模式Mode 子类Subclass 转换器Converter官方指南给出的整体工作流非常清晰所有 fake tensor 都关联一个FakeTensorMode。典型用法是先有一批真实张量分配一个FakeTensorMode用from_real_tensor把真实张量全部fake 化然后在 fake 张量上做分析。几个架构要点持久化的 memo 表保证别名一致性FakeTensorMode维护一张持久化的 memo 表将张量以及 storage映射到同一个 storage同一个张量多次 fake 化得到同一个fake tensor两个互相别名的张量 fake 化后得到两个共享同一个 fake storage的 fake tensor。这在源码中有明确对应FakeTensorConverter的tensor_memo属性torch/_subclasses/fake_tensor.py#L393-L400直接委托给self.meta_converter.tensor_memo——即MetaConverter内部的WeakValueDictionary以MetaTensorId为键缓存已转换的FakeTensor。是关系而非包装关系metadata 单一来源FakeTensor 被表示为一个__torch_dispatch__tensor subclass其底层元素是 meta tensor。也就是说fake tensor 底层就是 meta device 张量它通过额外的扩展钩子特别是dispatch_device来谎报真实的设备。这个设计刻意避免了维护两套元数据meta tensor 的元数据 fake tensor 的元数据的同步问题——is-a关系保证了只有一份权威元数据。源码佐证FakeTensor.device属性torch/_subclasses/fake_tensor.py#L897-L903的实现逻辑是当fake_mode.in_kernel_invocation为真时返回torch.device(meta)否则返回self.fake_device。in_kernel_invocation标志torch/_subclasses/fake_tensor.py#L717 的in_kernel_invocation_manager上下文管理器负责在内核执行期间让所有检查如is_meta、设备分配按 meta 设备行为执行。警告segfault 排查的第一方向官方指南给出了一个极具实践价值的排错提示早期 FakeTensor 最容易出错的地方就是太擅长伪装——fake tensor 伪装成 CPU/CUDA 张量后CPU kernel 拿到了 fake tensor 并试图解引用数据指针显然会崩溃。如果你在 fake tensor 代码中遇到段错误第一件事就是检查 C backtrace 落在哪个 kernel是 CPU kernel不正常还是 meta kernel正常。meta kernel 与真实 kernel 类似但它只负责分配输出不做任何数据计算。通用 fake 化配方tensor subclass 必须定义如何实现各类算子操作。官方指南给出了一般配方在输入 fake tensor 上运行 meta kernel把它们重新解释为 meta tensor——通过in_kernel_invocation_manager这个神奇上下文管理器完成。它指示整个 PyTorch 把 fake tensor 视为底层 meta tensor而不是把 fake tensor解包成 meta tensorfake tensor 本来就是 meta tensor若是工厂函数factory function直接以devicemeta调用底层工厂函数把结果 meta tensor 转回 fake tensor并计算输出设备通常很平凡但有时不平凡例如 CPU 标量提升cpu scalar promotion、设备转换类算子。API 实战三种典型用法非 PT2 用法独立使用官方指南给出的最直接用法更多示例见 test/test_fake_tensor.py# 创建 fake mode from torch._subclasses.fake_tensor import FakeTensorMode fake_mode FakeTensorMode() converter fake_mode.fake_tensor_converter # 将一些真实张量 fake 化 fake_x converter.from_real_tensor(fake_mode, x) with fake_mode: # 在 fake 张量上执行算子 fake_y fake_x * 2 # 工厂操作在上下文管理器内自动被 fake 化 fake_z torch.empty(20)这里回答了一个自然疑问为什么输入必须是真实张量因为在 PT2 场景下通常是即时编译JIT——编译某个图时你已经持有该图输入对应的真实输入你正在执行程序的过程中进行编译。from_real_tensor的完整签名torch/_subclasses/fake_tensor.py#L487-L497比指南展示的更丰富还支持make_constant、shape_env、source、symbolic_context、trace等参数用于在 Dynamo/AOTAutograd 上下文中传递符号上下文source会被用于构造StatefulSymbolicContext。PT2 pre-AOTAutograd 用法少见# Fake mode 尚未启用 from torch._guards import detect_fake_mode fake_mode detect_fake_mode(args) # 若 fake_mode 不为 None converter fake_mode.fake_tensor_converter fake_args [converter.from_real_tensor(fake_mode, arg) for arg in args] with fake_mode: ... # 用 fake args 做一些事情如需要detect_fake_mode会在多个位置搜索那个与生命周期关联的 fake tensor mode——通常是从 tracing context 中取出。源码印证torch/_guards.py#L1529-L1553若TracingContext中已存在fake_mode优先使用它这是 Dynamo 驱动编译时的情形否则按优先级启发式探测当前 dispatch mode 栈上激活的 fake mode → 传入张量关联的 fake mode。PT2 post-AOTAutograd 用法# Fake mode 已启用example_inputs 通常已经是 fake 的 # TODO: 我们可能想改变这一点 # 仍然通过它来访问 fake mode fake_mode detect_fake_mode(example_inputs) # 但通常你不必再显式开启它临时禁用 fake modefrom torch._subclasses.fake_tensor import unset_fake_temporarily with unset_fake_temporarily(): ... # 此处 fake mode 被禁用可以做真实张量计算何时需要禁用官方指南明确通常不需要。一个验证过的 niche 场景是在 fake tensor 上实现常量传播constant propagation——这时需要在 fake mode 中做一点真实的张量计算。源码中该工具定义于 torch/_subclasses/fake_tensor.py#L214返回一个可恢复先前 fake 状态的生成器。在 FX 图上传播 fake tensorFakeTensorPropimport FakeTensorProp from torch.fx.passes.fake_tensor_prop # 注应为 from torch.fx.passes.fake_tensor_prop import FakeTensorProp gm: GraphModule real_inputs: List[Tensor] FakeTensorProp(gm).propagate(*real_inputs) # 这会用 fake tensor 填充所有 FX 节点的 meta[val] # 若已有现成 fake mode应当复用它 FakeTensorProp(gm, modefake_mode).propagate(*real_inputs) # 若输入已经是 fake 的使用 propagate_dont_convert_inputs fake_inputs: List[FakeTensor] FakeTensorProp(gm, modefake_mode).propagate_dont_convert_inputs(*fake_inputs)源码实现torch/fx/passes/fake_tensor_prop.py#L16-L28解释了这个 API 的定位它继承torch.fx.Interpreter逐节点执行 FX 图并记录代表该节点元数据的 fake tensor。与ShapeProp相比它有两个优势便宜用 meta tensor 传播不真正存储数据且信息粒度更细例如通过 storage 查询准确的别名信息。实现细节也值得注意propagatetorch/fx/passes/fake_tensor_prop.py#L104-L109会把真实输入通过self._mode.from_tensor(a)全部转成 fake 后委托给propagate_dont_convert_inputspropagate_dont_convert_inputstorch/fx/passes/fake_tensor_prop.py#L111-L125在with self._mode:中运行整个图并保存/恢复输入 fake tensor 的 fake_device避免就地算子如shallow_copy_data_在传播期间永久污染调用方的输入每个节点运行后调用rebind_unbacked(shape_env, n, result)处理 unbacked 符号并把结果写入n.meta[val]同时若有写入n.meta[unbacked_bindings]。细节自动转换与元数据突变自动 fake 化还是报错最初FakeTensorMode不会自动 fake 化真实张量——如果你在 FakeTensorMode 区域内对真实张量做计算会直接报错。动机是防止这个 footgunwith FakeTensorMode(): real_tensor.t_()这段代码应该做什么如果它真的修改了真实张量的元数据会非常令人意外但同时又没有明显的契机去创建 FakeTensor。于是保守决策是抛错Invoking operators with non-Fake Tensor inputs in FakeTensorMode is not yet supported. Please convert all Tensors to FakeTensors first.这个错误在实践中相当烦人。典型场景你有一个真实的nn.Module想喂 fake 张量通过它——你得先想办法把整个nn.Modulefake 化。这直接催生了FakeCopyModetorch/_subclasses/fake_tensor.py#L3657一个TorchFunctionMode子类。最终官方放弃抵抗加入了自动 fake 化。但注意在FakeTensorMode的许多使用场景中自动 fake 化至今仍未默认开启——对应构造参数即allow_non_fake_inputs: bool Falsetorch/_subclasses/fake_tensor.py#L1556、torch/_subclasses/fake_tensor.py#L1606-L1608它控制是否允许在混合真实权重/全局变量与 fake 输入上调用算子。元数据突变把 fake tensor 存进 meta[val] 的隐患如果你持有 fake tensor 并执行t_()fake tensor 上的元数据会改变。表面合理但当你把 fake tensor 作为元数据存在 FX 节点上时突变会令旧元数据失效官方指南坦诚地指出了这里的根本矛盾fake tensor 维护极其精确的元数据精确到对象同一性object identity。而如果对象元数据在 FX 图生命周期内随时间变化图上其实没有任何方式表示这种随时间的变化。大多数严肃的 FX 分析都运行在函数化functionalized图上不存在这个问题但偶尔需要在非函数化图上做分析。官方甚至自嘲也许把 fake tensor 放进meta[val]本身就是个错误。About the tensor subclass双重分发模式FakeTensor 同时使用了subclass 模式 mode 模式两种 tensor subclass 机制FakeTensor.__torch_dispatch__会启用与该 fake tensor 关联的FakeTensorMode然后重新分发redispatch把繁重工作全部委托给FakeTensorMode如果 fake tensor 算子遇到不认识的 subclass 参数会返回NotImplemented把执行机会让给另一个 subclass 先运行希望它能把操作脱糖desugar成普通张量算子然后再次尝试。官方指南特别提醒这种机制可能引起无限循环——这是处理 subclass 交互时需要注意的已知风险。从FakeTensorMode的源码看它继承自TorchDispatchModetorch/_subclasses/fake_tensor.py#L1528并设置了_mode_key torch._C._TorchDispatchModeKey.FAKE向 torch_dispatch 分发基础设施表明自己是infra模式、分发优先级更低。构造函数还初始化了 dispatch 缓存统计、in_kernel_invocation标志、enter_stack记录进入模式前的栈用于退出时恢复先前的 fake mode等内部状态。每个算子是如何实现的官方指南指出任意给定算子的实现位置相当复杂需要了解的重要情形包括有限常量传播tensor subclass 支持在元素数量极小时做有限的常量传播帮助处理某些需要立即对这类张量调用item()的情况快速路径fastpath某些算子出于性能原因在 fake tensor 内部完全由手工实现传播规则custom_op如果使用custom_op生成自定义张量算子这些算子会直接向 fake tensor 注册impl_abstract设备转换算子的硬编码特例FakeTensor 本身对一些设备转换类算子有硬编码处理最后手段——真实执行如果没有 meta 实现也没有分解decomposition会生成真实的全零张量并直接运行该算子以观察结果。这可能在算子尝试用数据做索引时导致段错误因此对 custom op 默认不开启对应FakeTensorMode构造参数allow_fallback_kernels: bool True见 torch/_subclasses/fake_tensor.py#L1555。从源码还可以看到 fake tensor 对数据依赖类算子的特别照顾FakeTensor类上定义了nonzero_memo、item_memo、unique_memo、unique_consecutive_memo等SymNumberMemoDescriptor成员torch/_subclasses/fake_tensor.py#L851-L861用于对nonzero()、item()这类产生 unbacked SymInt 的算子做记忆化memoization。Converter 是如何工作的由于 fake tensor 用在对张量精确属性极其敏感的场景转换过程非常小心保留 leaf 性leaf-ness、requires_grad 属性、别名关系以及大量其他属性。繁重的工作主要由MetaConverter完成。源码印证torch/_subclasses/fake_tensor.py#L393-L408FakeTensorConverter在__init__中创建self.meta_converter MetaConverter(copy_datacopy_data)并提供add_constant_storage_mapping/invalidate_constant_aliases/clear_non_cpu_constants等方法管理常量constantfake tensor 及其别名——例如当const_tensor.add_(torch.rand([1]))修改了常量后该常量所有别名都必须不再为常量。转换入口from_real_tensortorch/_subclasses/fake_tensor.py#L487甚至允许你传入 meta tensor 来 fake 化虽然有点反直觉但在 cross ref 测试中会发生内部测试已经在 meta tensor 上操作了。如果已有 meta tensor也可以直接调用from_meta_and_devicetorch/_subclasses/fake_tensor.py#L681。性能特征Python 开销主导你可能会认为 fake tensor 不做任何计算所以很快——但实际恰恰相反在张量尺寸很小时我们完全被开销overhead主导。fake tensor 是 Python 实现的而且常常为单个张量算子做大量工作因为它们通过分解decompositions实现。所以 fake tensor 在实践中其实相当慢尤其是涉及符号形状symbolic shapes时。目前 fake tensor 中有两个重要的快速路径在实践中影响巨大Pointwise 算子不经过分解而是手工编写了它们的传播规则如果可能应当优先走 fastpath。此外源码中还内置了dispatch 缓存机制FakeTensorMode上有类级cache: dict[_DispatchCacheKey, _DispatchCacheEntry]、cache_hits、cache_misses、cache_bypasses统计以及epoch字段——每次用同一个 fake tensor mode 重新追踪时推进 epoch避免复用旧的 unbacked memotorch/_subclasses/fake_tensor.py#L1529-L1535。缓存开关由torch._dynamo.config.fake_tensor_cache_enabled与fake_tensor_cache_crosscheck_enabled控制。Fake tensor of fake tensor社区有兴趣把 fake tensor 作为用户输入直接送进 PT2 栈这暗示需要支持fake tensor 的 fake tensor。官方指南的结论是目前并未真正支持但实现起来可能不会太困难。与动态形状Dynamic Shapes的交互每个FakeTensorMode都包含一个ShapeEnv后者追踪所有符号形状信息二者的生命周期通常绑定在一起——同生共死。因为FakeTensorMode拥有ShapeEnv而 meta 实现没有数据依赖且需要分配 unbacked SymInt 的 meta 函数就落在 fake tensor 层实现。FakeTensor 还负责对 unbacked SymInt 做 memoization例如对同一个 fake tensor 调用两次nonzero()会得到相同的符号尺寸。这对应前面提到的nonzero_memo等成员。源码佐证FakeTensorMode.__init__接受shape_env: ShapeEnv | None None参数且static_shapes在未显式指定时默认为shape_env is Nonetorch/_subclasses/fake_tensor.py#L1557、torch/_subclasses/fake_tensor.py#L1585-L1588——即没有 ShapeEnv 就意味着静态形状。官方指南还提到双重 fake mode场景torch/_subclasses/fake_tensor.py#L1648-L1657 的注释导出export一个 fake model 时用户创建的外层 fake mode 通常没有 ShapeEnv而 Dynamo 创建的内层 fake mode 会持有 ShapeEnv 并替换为符号尺寸——因此仅对单个 FakeTensor 做isinstance测试是不够的。总结什么时候用 FakeTensor综合官方指南与源码FakeTensor 的适用场景可以归纳为Dynamo 追踪回答中间张量的 shape/dtype/device 查询支持数据依赖分支编译器 pass 元数据把 fake tensor 挂到 FX 节点的meta[val]注意函数化图上更安全常量传播结合unset_fake_temporarily在 fake 域内做少量真实计算动态形状分析通过FakeTensorMode内建的ShapeEnv传播符号尺寸并 memoize unbacked SymInt批量大小估算等分析类任务官方Other resources中提供了使用 FakeTensor 确定最大 batch size 的 Colab 教程链接社区中也有大量此类实践。需要记住的约束同样明确fake tensor 是 Python 实现的、小尺寸下开销主导对非 fake 输入默认不自动转换算子实现散落在 meta kernel、分解、fastpath、impl_abstract与 fallback 真实执行等多处以及最经典的排查建议——遇到段错误先看 C backtrace 落在 CPU kernel异常还是 meta kernel正常。这些细节构成了理解 PyTorch 2.x 编译流水线不可或缺的一块拼图。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考