Pyrefly 将张量形状引入 Python 类型系统:符号算术、Dim 与 DSL 驱动的 PyTorch 形状静态检查
Pyrefly 将张量形状引入 Python 类型系统符号算术、Dim 与 DSL 驱动的 PyTorch 形状静态检查【免费下载链接】pyreflyA fast type checker and language server for Python项目地址: https://gitcode.com/GitHub_Trending/py/pyrefly本篇文章基于 Pyrefly 团队在 PyCon US 2026 Typing Summit 上的演讲「Tensor Shapes in the Type System」整理展开系统讲解 Pyrefly 如何把张量形状tensor shapes变成 Python 类型系统的一等公民模型作者只需标注少量类/函数边界的类型参数所有中间张量的形状即可被自动推断并以行内类型提示inlay type hint呈现从而把「手写形状注释」变成「类型检查器生成的类型」。读完本文你将理解支撑这一能力的三大设计支柱类型级符号算术、Tensor/Dim两个用户可见类型、描述算子形状行为的微型 Python DSL掌握其工作原理、覆盖率边界、走向生产使用的现实约束以及如何在本仓库中快速上手体验。NanoGPT forward 未开启形状跟踪NanoGPT forward 开启形状跟踪问题的提出为什么张量形状不是类型系统的一部分演讲从一个很直白的观察切入一个 PyTorch 模型本质上就是一个「从张量到张量的函数」。中间是一连串被称为 tensor ops 的变换把输入张量逐步变成输出张量——你可以把它看作一条直线型程序或者一张计算图。而模型开发中最难的部分恰恰不是架构设计而是组合算子时跟踪形状。就像普通编程中你要考虑调用什么函数、如何组合它们的类型一样在 PyTorch 程序或者说模型里形状就是类型的单位。当下的标准做法是在代码里写大量注释来记录形状。因为这些算子都是非平凡的复杂变换开发者只能一边写一边用注释记下「这一行张量是什么形状」。演讲中以 Andrej Karpathy 的 nanoGPT 为例——一个教学性质的微型 GPT-2 实现——展示了这种惯例创建pos张量时注释形状T经过一堆 transformer 调用后进入循环逐个调用 block 模块而每个 block 内部还有自己的张量算子贯穿始终作者在每行写下B, T、embedding 等形如元组的形状说明。演讲提出的核心问题是如果像普通类型那样类型检查器直接给出类型提示开发者是不是就不用再写这些注释了答案正是 Pyrefly 的张量形状功能——同样的代码函数边界照常做类型标注这是普遍实践配合合适的泛型所有局部变量的类型都会被自动推断形状在行内类型提示中对齐呈现。这不是魔法你标注的是函数边界而所有中间结果的类型由检查器推得。三大设计支柱Pyrefly 是如何做到的演讲将实现思路归纳为三个独立的设计决策每一个都可以独立理解。支柱一类型层面的符号算术第一个想法与张量本身无关而是在类型系统层面引入符号算术。泛型类型本来就可以接受类型参数这次的新颖之处在于这个类型参数表示的是维度大小——它可以是普通整数可以是符号也可以是这些符号组成的算术表达式。它们本质上都属于int但是符号化表达式。这里有一个关键的设计取舍不把 SAT 求解器塞进类型检查器。系统不会去问「是否存在某个 X 使得 13 X 6」这类可满足性问题。取而代之的是符号表达式之间的相等判定通过规范化normalization过程完成——把两侧表达式化简然后做语法相等比较。规范化过程内置了一批应当成立的等式规则尽力而为。演讲明确承认这并不保证所有问题可判定但在实践中表现良好而且系统不引入存在类型existential types这就是这套方案的边界所在。在仓库中类型级算术的化简逻辑可以从 pyrefly_types/src/simplify.rs表达式化简与 pyrefly_types/src/lit_int.rs字面量整数表示等模块的源码结构中得到印证。支柱二仅两个用户可见类型——Tensor与Dim有了类型级符号算术之后系统只向用户暴露两个类型。第一个是张量类型Tensor它以维度作为类型参数。任意维度数量都合法因为张量本身就是多维的当维度数量未知或某些维度未知时随时可以退化为Any——这就像普通的可选类型optional typing一样宽松。第二个是Dim。形状必须有一个来源所有运算的根基是一些张量创建算子它们接收整数作为参数也就是模型里的配置数值由这些数值决定权重矩阵的大小。一切最终归结为整数而Dim就是用来把符号算术包装起来、让整数流动到需要它们的地方的类型。为什么是Dim演讲给出的心智模型非常精妙Dim之于Literal就像符号整数之于具体整数。Literal只接受具体整数而Dim同时接受具体整数和符号整数。如果 Python 类型系统里已经有包装符号的东西直接用它就好了——所以Dim是一个非常轻量的补充本身不是系统的核心但没有它整数就无法被传播到目标位置。在这个体系中Dim在用户代码里写作Int[X]由shape_extensions包导出负责把运行时整数值桥接到类型级符号。比如x: Tensor[[3, 4]]时x.shape的类型是tuple[Int[3], Int[4]]你可以从张量中取出维度并用它构造新张量算术同样成立a: Int[3]与b: Int[4]相乘得到Int[12]。泛型类型参数则让模块具备形状多态性官方文档 tensor-shapes.mdx 给出了一个典型例子class Linear[N: IntVar, M: IntVar]: def __init__(self, n: Int[N], m: Int[M]): ... def forwardXs: IntTuple - Tensor[[*Elements[Xs], M]]: ... linear: Linear[3, 4] Linear(3, 4) inp: Tensor[[2, 5, 3]] ... x: Tensor[[2, 5, 4]] linear(inp)Int在整个 PyTorch 生态中编码符号形状不仅作用于张量也作用于模块如nn.Linear[3, 4]。类型变量上的算术还允许编写自定义形状变换def custom_rand_tensorA: IntVar, B: IntVar - Tensor[[(A B) // 2]]: return torch.randn((a b) // 2) x: Tensor[[3]] custom_rand_tensor(2, 4)支柱三用于算子形状行为的微型 Python DSL第三个想法在工程实现上最为关键。PyTorch 有成千上万个算子它们的形状变换非常花哨——不是难以理解而是很难用类型系统表达。有些算子的形状签名很简单用标准类型桩type stubs就能描述。例如矩阵乘法mm一个n×k张量与k×m张量相乘得到n×m对所有类型参数实例化都成立def mmM: IntVar, K: IntVar, N: IntVar - Tensor[[M, N]]: ...但更多算子的形状变换与其不断提高类型系统的表达能力去硬编码不如程序化地描述。演讲指出这正是此前一些尝试的瓶颈所在想把算子编码在类型层面到某个点就不得不放弃。Pyrefly 的做法深受PyTorch 编译器自身的 symbolic shapes符号形状实现启发为每个算子提供一个「fake op」。fake op 模拟原始算子的行为但只关注形状层面、不涉及数据——它只回答一个问题给定这种形状的张量输出张量是什么形状以repeat算子为例它把张量沿若干维度按指定次数重复。用一行 Python 推导式就能表达其形状行为。因此 Pyrefly 使用了一个极小的 Python 子集——可以把它理解为「对整数做列表推导」——事实证明这足以覆盖成千上万个算子中的绝大部分。这些算子都以特定方式声明而少数算子如矩阵乘法根本不需要这种机制普通类型签名就够了。有了这些声明再用符号算术把它们贯通起来就能推导出类型。仓库中的具体实现位于 tensor-shapes/pyrefly-torch-stubs/torch-stubs/_shapes.pyi官方文档 tensor-shapes.mdx 展示了这类 DSL 声明的样子内部库定义非用户可见代码type_shape_dsl_function def repeat_shape(shape: IntTuple, repeats: IntTuple) - IntTuple: if len(repeats) len(shape): return dsl.Invalid( Number of dimensions of repeat dims can not be smaller than number of dimensions of tensor ) extra len(repeats) - len(shape) return dsl.IntTuple( ( repeats[index] if index extra else shape[index - extra] * repeats[index] for index in range(len(repeats)) ) )这套设计的直接收益是扩展新 PyTorch 算子的形状覆盖无需触碰 Pyrefly 内部。贡献者只需添加 fixture 桩、DSL 函数或移植模型详见 tensor-shapes-contributing.mdx 与仓库根目录的 TENSOR_SHAPES_CONTRIBUTING.md。端到端效果从 MLP 到注意力机制这套系统在小模型上效果很好演讲给出了两类代表性例子。**多层感知机MLP**是最简单的构建块一个包含少量张量变换的小类。可以看到代码里有算术运算例如4 * config.n_embedding而这些数值会出现在类型里——形状以类型的形式可用。这个类只需要声明一个泛型参数n_embedding。仓库配套教程 tensor-shapes-tutorial-basics.mdx 用一个强化学习的BaselineActor三个nn.Linear串联完整演示了移植过程构造参数state_size、action_size决定张量维度因此必须是Int[S]、Int[A]并绑定到类级类型参数batch 维度每次调用都会变化因此作为方法级类型参数forward[B: IntVar]随后用assert_type逐个验证中间形状确认无误后移除每个assert_type都对应 IDE 中永久显示的行内类型提示class BaselineActorS: IntVar, A: IntVar: def __init__(self, state_size: Int[S], action_size: Int[A]) - None: super().__init__() self.fc1 nn.Linear(state_size, 400) self.fc2 nn.Linear(400, 400) self.out nn.Linear(400, action_size) def forwardB: IntVar - Tensor[[B, A]]: h1 F.relu(self.fc1(state)) # pyrefly infers: Tensor[[B, 400]] h2 F.relu(self.fc2(h1)) # pyrefly infers: Tensor[[B, 400]] act torch.tanh(self.out(h2)) # pyrefly infers: Tensor[[B, A]] return act当用户写下BaselineActor(24, 4)时类型检查器绑定S 24、A 4推断出类型BaselineActor[24, 4]子模块自动获得类型self.fc1是Linear[24, 400]self.out是Linear[400, 4]。**注意力块attention block**则更复杂一些self.c_attn是一组权重其第二个维度是三倍的嵌入数随后在该维度上做三分割得到三个张量的元组——此时普通类型机制接管接着是一串 transpose、floor div 等操作。演讲强调这些全部被支持「而且都顺理成章地工作」。标准注意力实现一堆矩阵乘法、masking 等在每一行都能得到类型提示Karpathy 原本为了让代码可读而手写的形状注释现在由类型提示自动完成。移植工作并不止于 nanoGPT。演讲提到团队还移植了一大批其他开源模型覆盖从现代到稍旧的各种架构以及 PyTorch 惯用的各类编码模式其中绝大多数都能顺利通过——这本身也测试了算子的覆盖广度。从仓库目录结构看这批移植模型与配套测试分布在 tensor-shapes/pyrefly-torch-stubs、tensor-shapes/pyrefly-numpy-stubs 等包的test与examples目录下。此外仓库还提供了可在线试玩的 sandbox 示例例如 tensor-shapes-overview/sandbox.py 用assert_type演示了形状推断、变长 batch 维variadic batch dims与 transpose 维度重排并包含一个故意写错返回形状的broken函数供你观察报错。覆盖率观察损失来源与有趣的边缘情况演讲对覆盖率给出了一些总体评论张量程序乍看吓人但归根结底是简单变换的组合因此覆盖率往往很高平均而言只需要很少的抑制suppression。主要覆盖损失来源与普通 Python 程序相同——例如包含非齐次元素的列表会导致类型信息丢失而标准的规避手段同样适用。这与张量检查本身无关。演讲还分享了两类反复出现、值得后续讨论的现象位置相关的类型有时你会有一个列表其中位置i的元素类型依赖i——类型是i的函数。如果类型系统能表达这一点会很理想。虽然可以通过在程序里做归纳induction这类高级技巧「勉强工作」但演讲者明确表示不指望普通程序员用归纳思维写代码。所以典型做法是进入时丢失形状通过注释在出口处恢复形状。整除信息声明符号时没有附带任何额外信息。但如果声明时能附带整除性信息就能证明更多等式——而这类等式在实践中反复出现。目前系统选择忽略它们还没有处理但如果在创建处多携带一点信息这些等式是可以被证明的。针对这些真实世界的缺口官方参考文档 tensor-shapes-reference.mdx 给出了一个实用的注解优先级从最可取到最后的兜底assert_type验证检查器推断证明系统在起作用→ 注解回退检查器无法推断但注解兼容需说明原因→type: ignore检查器产生了错误类型例如代数缺口必须注释说明具体缺口→ 裸Tensor形状确实不可知需说明具体原因。用 AI 来移植模型顺应时代主题这个项目大量使用了 AI——主要用于移植模型。nanoGPT 是手工完成的但移植过程非常机械天然适合用 AI agent 来做省去逐行手写。演讲者分享了一个有意思的观察LLM 经常惊讶于这类事情居然能在类型系统层面完成——这并不奇怪因为对类型检查器来说这是一个全新的领域。LLM 常常做出悲观的假设演讲者必须说「不试试看这会成功的」然后它试完会说「哇好吧我可真聪明」。进一步的自动化是技能skills把一套指令或工作流交给 agent 循环执行。仓库中就带有移植模型的技能 add-shape-types-to-torch-model内含shape_tracking_capabilities.md、style_guide.md等指导文件。团队试验过这套流程几乎一次one-shot就移植了几个模型效果良好。技能内置了一个循环LLM 移植完模型后会就类型覆盖率等指标自我评分、参考其他模型再决定是否继续。大多数情况两次尝试就够有时一次就成功。这正是演讲者想通过所有实验证明的核心观点这套系统在实践中可用而且对 AI 友好。走向生产使用当前的现实约束那么真正的阻碍在哪里演讲者「小小地作弊」了一下使用了 TypeVar 及其上的算术。问题在于如果在运行时求值这些注解会失败——原因很小但很致命Python 内置的typing.TypeVar不支持算术注解中的D // NHead这类表达式在注解被求值时抛出TypeError。解决方案有两条官方文档 tensor-shapes-setup.mdx 有完整说明from __future__ import annotations推荐把所有注解的求值推迟到运行时之后形状算术永远不会在运行时执行。它同时兼容旧式与新式泛型PEP 695 的class Foo[T]语法。shape_extensions.IntVar运行时兼容直接导入shape_extensions包它补丁了torch.Tensor、nn.Conv2d等 torch 类使其在运行时接受下标语法而不崩溃并提供一个支持算术的IntVarN 1返回自身而不是抛TypeError。但注意 PEP 695 新式泛型内部硬编码使用typing.TypeVar不允许算术所以此方案必须搭配旧式泛型。演讲还提到与同类系统的对比。最流行的大概是jaxtyping——它在运行时做这种检查有自己的一套语法。jaxtyping 的语法比 Pyrefly 提出的方案更重把维度放进字符串里如Shaped[Tensor, M 2 M//2]而不是Tensor[[M, 2, M // 2]]但确实是今天人们在使用的东西。Pyrefly接受 jaxtyping 注解语法jaxtyping true开启见 configuration.mdx目前尚未提供运行时类型检查能力因此互操作不是双向的。其局限在于jaxtyping 无法在类的多个方法和变量间共享符号维度这使其只能服务于操作单个张量的函数而无法像 Pyrefly 这样端到端地贯通整个模块层级——演讲者明确指出NanoGPT 这类真实模型的完整类型检查实现无法仅用 jaxtyping 语法忠实地移植。面向未来演讲提出两点期望放宽运行时限制如果把这项能力交到模型作者手中他们会发现非常易用。作为 Python typing 社区如果能推动取消 TypeVar 算术限制会很有帮助——语法对采纳至关重要。Dim的通用化符号算术部分与张量无关可能惠及更广泛的用户——不只 NumPy任何 SymPy 用户都可能受益。如果能去掉Dim这个额外的用户可见内建类型换成更通用的机制会更好。如何上手体验该功能在 Pyrefly 中目前是实验性的API 与行为可能在后续版本无通知地变化且目前只支持 PyTorch 张量演讲确认计划在不久的将来支持其他数据类型如 NumPy 数组。按 tensor-shapes-setup.mdx 的指引上手只需三步安装形状感知的桩pip install pyrefly-torch-stubs这是一个 PEP 561 stub-only 包只携带 PyTorch 的类型信息不碰运行时torch包。PyTorch 自带的类型桩不含形状信息因此这些桩会优先于它们提供形状感知版本例如nn.Conv2d.__init__把 kernel size、stride、padding 捕获为类型级数值forward计算输出空间维度。安装它同时会拉入pyrefly-shape-extensions提供shape_extensions包导出Int等。两个包与 Pyrefly 锁步版本化。当 Pyrefly 能解析到shape_extensions包时张量形状支持自动启用无需其他配置。可选使用仓库内的本地副本桩源码就位于仓库 tensor-shapes 目录下。把该目录拷进项目在pyrefly.toml中指向两个包目录search-path [ tensor-shapes/pyrefly-torch-stubs, tensor-shapes/pyrefly-shape-extensions, ]注意桩使用了 PEP 696 类型参数默认值class Conv1d[..., S: IntVar 1]Pyrefly 只在python-version为 3.13 及以上时解析该语法。要么把python-version调高它只影响检查所针对的版本不改变代码实际运行的解释器要么把副本从检查中排除导入仍会通过search-path解析python-version 3.13 # 或者 project-excludes [tensor-shapes/**]写一个 hello world 并检查创建一个hello_shapes.py用from __future__ import annotations、导入shape_extensions的Int/IntVar写一个带形状泛型的二层网络然后运行pyrefly check hello_shapes.py。无报错即说明形状已被推断在 IDE 中会看到h的形状以行内类型提示显示为Tensor[[B, HidDim]]。想深入可继续阅读仓库内的整套官方文档tensor-shapes.mdx总览、tensor-shapes-setup.mdx安装与配置、tensor-shapes-tutorial-basics.mdx、tensor-shapes-tutorial-loops.mdx循环与堆叠、多 head 注意力、tensor-shapes-tutorial-architectures.mdx编码器-解码器与指数形状的递归链、tensor-shapes-reference.mdxInt/Tensor/assert_type/reveal_type的完整 API 参考与 jaxtyping 兼容性说明。结语从 nanoGPT 的注释惯例出发Pyrefly 用「类型级符号算术 两个用户可见类型 微型 DSL」三个设计决策把张量形状变成了 Python 类型系统的一部分函数边界做少量标注中间形状全部自动推断在 IDE 中以行内类型提示呈现。它刻意避开了 SAT 求解器与存在类型把描述 PyTorch 数千算子的复杂度隔离到库维护者可扩展的 DSL 中用 AI 辅助完成了大量真实模型的移植验证。虽然仍受限于运行时 TypeVar 算术与实验性状态但正如演讲所强调的这套系统在实践中可用而且对 AI 友好——这正是它走向模型作者日常工具链的第一步。【免费下载链接】pyreflyA fast type checker and language server for Python项目地址: https://gitcode.com/GitHub_Trending/py/pyrefly创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考