Pyrefly 张量形状类型系统贡献指南:Fixture Stubs、类型级 Shape DSL 与验证工作流
【免费下载链接】pyreflyA fast type checker and language server for Python项目地址: https://gitcode.com/GitHub_Trending/py/pyrefly
Pyrefly 的张量形状(tensor shape)类型追踪系统被刻意设计为一个"可扩展的三层架构":绝大多数 PyTorch 形状覆盖可以通过编辑.pyistub 和测试完成,而无需触碰 Pyrefly 的 Rust 内核。本文基于仓库根目录的 TENSOR_SHAPES_CONTRIBUTING.md 展开,完整覆盖三类形状机制(Fixture stubs、类型级形状 DSL、Special handlers)的工作原理与编写规范,并给出从run_pyrefly.py静态校验、run_tests.py全量测试到cargo test shape_dsl内核测试的完整验证工作流。读完本文,你可以独立地为 Pyrefly 添加新的形状 stub、编写类型级 DSL 函数,并为移植的 PyTorch 模型补充assert_type形状检查点。
架构总览:三种互补的形状追踪机制
Pyrefly 的形状追踪由三种互补机制组成(详见 TENSOR_SHAPES_CONTRIBUTING.md):
- Fixture stubs(夹具 stub):带形状泛型签名的
.pyi文件,覆盖nn.Linear、nn.Conv2d这样的模块和torch.mm这样的函数。 - 类型级形状 DSL 函数:用一小套 Python 子集在 tensor-shapes/pyrefly-torch-stubs/torch-stubs/_shapes.pyi 中书写的形状变换,函数以
@type_shape_dsl_function装饰,并被公共返回注解直接调用。用于覆盖 reduction、padding、pooling、convolution 这类"需要计算"的形状逻辑。 - Special handlers(特殊处理器):Pyrefly 实现层的逻辑,服务于需要更深类型系统集成的一级语法,例如
nn.Sequential链式调用、.shape、.size()、assert_shape与装饰器解释。
前两种机制都位于tensor-shapes/目录下,是添加或改进形状覆盖的常规途径。发布版 stub 只使用类型级 DSL;Pyrefly 目前仍保留旧的@shape_dsl_function与@uses_shape_dsl(...)(V1 机制)的内核支持与隔离测试,仅为了让已锁定的 V1 stub 在迁移期间保持兼容——贡献者不应再添加新的 V1 规则。Special handlers 则需要改动 Pyrefly 实现,应按"内核工作(kernel work)"对待。
从源码结构看,运行时支撑来自 tensor-shapes/pyrefly-shape-extensions/shape_extensions/init.py:这个 Python 包提供Int、IntTuple、IntVar、D、assert_shape、type_shape_dsl_function等原语的最小运行时类——正如文件头注释所述,.pyistub 为 Pyrefly 提供完整类型信息,而.py文件只保证"这些注解在 Python 中求值时不会崩溃"。它还包含一个实用的兼容性技巧:当检测到 torch/jax 已安装时,会为torch.Tensor、nn.Linear等类动态打上__class_getitem__,使Tensor[B, T, N]、nn.Linear[In, Out]这样的下标注解在运行时成为无操作(no-op)而不是抛出type is not subscriptable异常。
Fixture Stubs:用形状泛型签名描述张量变换
存放位置
tensor-shapes/pyrefly-torch-stubs/torch-stubs/ |-- __init__.pyi |-- _shapes.pyi |-- nn/ | |-- __init__.pyi # nn.Linear, nn.Conv2d, nn.LSTM, etc. | `-- functional.pyi # F.relu, F.softmax, F.conv2d, etc. |-- distributions/ | `-- ... # torch.distributions `-- ...张量形状测试运行器会把tensor-shapes/作为 Pyrefly 搜索路径传入,因此这些 stub 在验证时会覆盖常规的torchstub。
Stub 是如何工作的
Fixture stub 提供"形状泛型"的类型签名。以nn.Linear为例:
class LinearN, M: def __init__( self, in_features: SymInt[N], out_features: SymInt[M], bias: bool = True, ) -> None: ... def forward*Xs -> Tensor[*Xs, M]: ...构造函数把输入/输出维度捕获为类型参数;forward方法再借助变长类型参数*Xs表示可任意透传的 batch 维度。仓库中的实际实现 tensor-shapes/pyrefly-torch-stubs/torch-stubs/nn/init.pyi 采用了当前推荐写法:类型参数以IntVar为界,forward用IntTuple+Elements解构批量维度:
class LinearIN: IntVar, OUT: IntVar: weight: Tensor[[OUT, IN]] bias: Tensor[[OUT]] | None def __init__( self, in_features: _Int[IN], out_features: _Int[OUT], bias: bool = True, ... ) -> None: ... def forwardBs: IntTuple -> Tensor[[*Elements[Bs], OUT]]: ...可以看到构造参数通过_Int[IN]/_Int[OUT](即Int[...]包裹)与类型参数绑定,而forward中*Elements[Bs]把任意 rank 的批量维度解包后原样透传——这正是"batch 维度不变、仅最后一维从IN变为OUT"这一形状语义的声明式表达。同文件中的Dropout、GELU则展示了最简形态:forwardShape: IntTuple -> Tensor[Shape],即形状完全透传。
编写新 Stub 的步骤
- 识别形状签名:输入维度、输出维度以及它们之间的关系。
- 对"决定张量维度"的参数使用
SymInt[X](当前实现中等价于用IntVar界类型参数 +Int[...]注解);bias、dropout这类非形状参数保持原始类型。 - 写出表达形状变换的方法或函数签名。对原样透传的 batch 维度使用
*Xs或*Bs。 - 把 stub 加入 tensor-shapes/pyrefly-torch-stubs/torch-stubs 中对应的
.pyi文件。 - 在
tensor-shapes/pyrefly-torch-stubs/test/下添加或更新聚焦测试。
示例:添加一个新模块
假设要添加保持空间维度不变的nn.GroupNorm:
class GroupNormNumGroups, NumChannels: def __init__( self, num_groups: SymInt[NumGroups], num_channels: SymInt[NumChannels], eps: float = 1e-5, affine: bool = True, ) -> None: ... def forward*S -> Tensor[*S]: ...由于GroupNorm不改变形状,forward 签名就是简单的Tensor[*S] -> Tensor[*S]。
类型级 Shape DSL 函数:当签名不足以表达输出形状时
当一条普通签名无法表达输出形状时(例如 conv 的 stride/padding 计算、reshape 的-1推导),就需要类型级 DSL。
存放位置与调用方式
DSL 函数统一放在:
tensor-shapes/pyrefly-torch-stubs/torch-stubs/_shapes.pyi公共 stub 在返回注解中直接调用类型级 DSL 函数。例如:
from shape_extensions import IntTuple, type_shape_dsl_function import shape_extensions.dsl as dsl @type_shape_dsl_function def repeat_shape(shape: IntTuple, repeats: IntTuple) -> IntTuple: if len(repeats) < len(shape): return dsl.Invalid("repeat dimensions cannot be shorter than the input rank") extra = len(repeats) - len(shape) return dsl.IntTuple( repeats[i] if i < extra else shape[i - extra] * repeats[i] for i in range(len(repeats)) ) def repeatShape: IntTuple, Repeats: IntTuple -> Tensor[repeat_shape(Shape, Repeats)]: ...真实的 _shapes.pyi 已包含十余个这样的函数。以reduce_shape为例,它对dim is None(全归约)、is_int_value单维、负数维度、空元组、0 维标量的0/-1别名、重复维度等边界都做了显式处理,越界或重复维度统一返回dsl.Invalid("dimension out of range")/dsl.Invalid("duplicate dimension");reshape_shape则用dsl.is_concrete_int区分具体与符号量,在只允许一个-1的前提下推导其大小,并对元素总数不匹配给出dsl.Invalid诊断。这些实现印证了下面"DSL 子集"一节的能力清单。
DSL 子集:有意保持很小的语言
DSL 是有意的一个"小而代数化"的语言,其主要值域是Int(单个形状维度)与IntTuple(完整形状);运行时配置值通过公共签名上的Flag[...]类型参数接入。函数体支持:
dsl.IntTuple(...)构造结果形状;len、索引、切片,以及有界生成器表达式;+、-、*、//、%等算术运算;if/else;- 单赋值局部变量;
- 直接
return其他@type_shape_dsl_function帮助函数的结果; dsl.concat、dsl.prod、dsl.Invalid等 DSL 操作;dsl.Int.gradual()表示一个"渐进维度"表达式;return dsl.IntTuple.gradual()直接返回一个渐进形状。
DSL 函数应保持简单、代数化。它们由 Pyrefly 分析,而不是PyTorch 运算的常规运行时实现。
整型参数与 IntVar 的传递边界
这是 DSL 写作中最容易出错、也最"无法靠猜"的部分,规范如下:
- 帮助函数参数:消费单个维度时声明为
Int;None有独立含义时声明为Int | None。 - 公共签名中的运行时整数若需透传给帮助函数,优先用一个"界恰好为形状
Int"的类型参数,并用该类型参数标注运行时参数:
@type_shape_dsl_function def resize_shape(size: Int) -> IntTuple: return dsl.IntTuple((size,)) def resizeN: Int -> Tensor[resize_shape(N)]: ...这种形式要求 bound 恰好是Int。
- 当类型参数只是在直接
IntTuple或 list 形状语法中命名一个符号维度时,改用IntVar。若这个符号维度还要传给 DSL 帮助函数,必须在调用边界用Int[...]包裹:
def zerosN: IntVar -> Tensor[[N]]: ... def resize_symbolicN: IntVar -> Tensor[resize_shape(Int[N])]: ...- 裸
IntVar参数以及N + 1这类算术在帮助函数调用中会被拒绝,应写成Int[N]与Int[N] + 1。 - 直接写成实参的
Int[N] | None是类型联合语法,不是运行时 DSL 值;应按语义传Int[N]、None,或 bound 恰好为Int | None的类型参数。 - 仍解析为
Int | None的实参在控制流收敛前按渐进(gradual)接受,非None分支中可作为Int使用。 - 宽泛的运行时
int用作维度时会变成渐进维度,保留已知 rank 与其他维度;Any则保持未知,不会被当作渐进整数。 dsl.Int.gradual()本身是Int表达式,可参与算术与dsl.IntTuple(...)构造;dsl.IntTuple.gradual()目前只能作为 DSL 函数的直接返回表示整个渐进形状,不能赋值给局部变量或嵌入更大的表达式。D[...]与D(...)是兼容包装器,用于 Python 会急于求值的注解场景。它们不能替代Int[...]:D[N]里仍是裸IntVar,会被拒绝;D[Int[N] + 1]合法。- 当某个分支需要"在形状求值期已知的整数字面量"时,用
dsl.is_concrete_int(value)(接受Int或Int | None);它对None、符号维度、渐进Int均为False。用dsl.is_int_value(value)收窄兼容的Flag[int | tuple[int, ...] | None]值中的整型成员——注意它不能证明该整数是具体的。
关于 NumPy 和 JAX stub 所用的另一套更小的类型级 DSL 子集(仍在建设中),有两条"无法靠猜"的经验值得记住:
- 一个使用了不支持语法的 DSL 函数会在所有调用点求值为
Unknown,而调用点本身不报告任何错误。真正的方法是直接类型检查 stub 文件——测试运行器会以stubssuite 的形式替你完成这一步。 int | tuple[int, ...]类型的参数仅靠is_int_value收窄后不能迭代。以is None检查开头可以让收窄生效,因此这类参数应声明为int | tuple[int, ...] | None,函数体中拒绝None。Torch stub 的conv_shape与 JAX stub 的reshape_shape都采用这个写法。
示例:reduction
@type_shape_dsl_function def reduce_shape(shape: IntTuple, dim: int, keepdim: bool) -> IntTuple: axis = dim % len(shape) return dsl.IntTuple( 1 if keepdim and i == axis else shape[i] for i in range(len(shape)) if keepdim or i != axis )公共 stub 会把输入形状与运行时选项绑定到类型参数,然后在返回注解中调用reduce_shape(...)。
添加新 DSL 函数的步骤
- 在 tensor-shapes/pyrefly-torch-stubs/torch-stubs/_shapes.pyi 中书写形状变换。
- 用
@type_shape_dsl_function装饰。 - 用
Int、IntVar、IntTuple或Flag[...]类型参数绑定相关公共参数,并在返回注解中调用 DSL 函数。 - 添加使用
assert_type检查计算后形状的正向测试。 - 若 DSL 应当拒绝非法形状或报告形状错误,添加带
# E:期望的负向测试。
旧的基于装饰器的 DSL 仅保留给尚未迁移的规则。避免在新规则中混合 V1 与 V2 逻辑;如果 V2 目前无法表达某个操作,记录该缺口(document the gap),而不是增加新的 V1 表面积。
当前已知限制
类型级 DSL 使用一个小的、可组合的语言,并不会精确建模所有形状行为。当前已知限制包括:符号arange的取整、unfold与diag_embed的符号配置值、结构化tensordot轴列表、符号 rank 形状或派生符号维度的乘积、以及 list 形式的 padding。遇到这些情况应尽可能保持渐进(gradual),补充聚焦测试,并在受影响的规则处留下TODO(stroxler),使精度损失保持可见。
移植模型(Ported Models)
存放位置
tensor-shapes/pyrefly-torch-stubs/examples/每个文件都是对某个真实 PyTorch 模型的完整标注移植,带assert_type检查点与冒烟测试。
添加新模型的步骤
- 从 TorchBench(PyTorch 官方基准模型集)或其他来源挑选一个模型。
- 参照仓库内教程(tensor-shapes-tutorial-basics 文档)或 Agent 移植技能(tensor-shapes/skills/add-shape-types-to-torch-model)完成移植。
- 在形状变化操作之后添加
assert_type或assert_shape检查点。 - 运行时执行有价值时,在文件底部添加冒烟测试。
- 运行
verify_port.sh检查常见质量问题。
verify_port.sh检查项
该脚本检查移植模型中的常见问题:
tensor-shapes/skills/add-shape-types-to-torch-model/verify_port.sh tensor-shapes/pyrefly-torch-stubs/examples/<model>.py它报告如下指标:
| 指标 | 含义 |
|---|---|
ig | type: ignore计数 |
bs | 签名中的裸Tensor计数 |
bv | 变量注解中的裸Tensor计数 |
sh | 带形状的assert_type计数 |
ba | 裸assert_type计数 |
sm | 冒烟测试计数 |
测试 Stub 与示例变更:tensor-shape 专用 Pyrefly 运行器
对多数贡献而言,最重要的验证是 tensor-shape Pyrefly 运行器。它使用形状感知的 stub 检查聚焦测试、负向期望、jaxtyping 示例与示例语料库,并且还会类型检查 stub 文件本身(以stubssuite 形式报告)。
这一点比听起来更重要:Pyrefly 只对它被要求检查的文件报告错误,因此通过--search-path触达的 stub 是"沉默的"。一个无法编译的 stub 不会自我暴露,它只是停止贡献类型,让所有调用点安静地推断出Unknown——这看起来像"规则缺失"而不是"规则损坏"。直接检查 stub 能把这种情况变成带行号的错误。
Torch 包目前通过其 run_pyrefly.py 中的check_stubs=False暂时退出该检查(源码中对应的 TODO 指出:包中仍存在未完成的内部导入、类型参数遮蔽等问题,且 torch-stubs/_shapes.pyi 的 V1@shape_dsl_function函数体不是合法 Python);类型级 DSL 文件可以干净通过检查,因此迁移这些规则到类型级 DSL 正是移除该退出项的方式。
基本用法
python3 tensor-shapes/pyrefly-torch-stubs/run_pyrefly.py运行器会自行构建 Pyrefly——默认用cargo build,加--buck时用 Buck——因此始终针对你的工作副本检查。先单独构建一次是多余的;跳过构建则意味着运行结果来自旧版 Pyrefly。
从 run_pyrefly.py 的参数定义可以确认完整的命令行选项:
--pyrefly /path/to/pyrefly:显式传入二进制的唯一不构建模式——因为裸路径无法表达如何重新构建它;--buck:改用 Buck 构建并运行;--release:用 Cargo release profile 而非 debug 构建;--python:指定提供 Torch fallback 模块的解释器(默认共享 virtualenv);--suite:可重复,只运行指定 suite,默认全部;--nocapture:流式打印 Pyrefly 完整输出(默认只在失败时转储检查器输出,成功时打印紧凑的PASS ...行)。
构建使用自定义目标目录时,run_pyrefly.py会遵守CARGO_TARGET_DIR。
迭代时单跑某个 suite
suites.py 定义了 Torch 包的五个 suite,--suite的可选值正来源于此:
| Suite | 匹配文件 | 说明 |
|---|---|---|
torch-examples | examples/*.py、examples/runtime/*.py | 示例语料库 |
torch-positive | test/test_*.py | 正向测试 |
torch-negative | test/negative_tests/test_*.py | 带# E:期望的负向测试 |
jaxtyping-positive | test/jaxtyping/test_*.py | jaxtyping 集成(Python 3.12 配置) |
jaxtyping-negative | test/jaxtyping/negative_tests/test_*.py | jaxtyping 负向测试 |
python3 tensor-shapes/pyrefly-torch-stubs/run_pyrefly.py --suite torch-positive python3 tensor-shapes/pyrefly-torch-stubs/run_pyrefly.py --suite torch-negative python3 tensor-shapes/pyrefly-torch-stubs/run_pyrefly.py --suite torch-examplesstub 没有对应的 Buck 测试目标;内部 checkout 通过--buck运行同一运行器,只是以不同方式获取 Pyrefly:
python3 tensor-shapes/pyrefly-torch-stubs/run_pyrefly.py --buck一次运行所有库(与 CI 完全一致)
python3 tensor-shapes/run_tests.py # 内部 checkout 追加 --buck python3 tensor-shapes/run_tests.py --static-only由于 Torch 与 NumPy 的形状 stub 会回退到已安装库中的定义,即使--static-only也需要共享 virtualenv。可用--python选择安装了所需库的另一个 virtualenv 解释器;根运行器把它用于运行时测试,并转发给 Torch 与 NumPy 的静态检查。
项目的 test.py 运行器把张量形状验证与默认 Pyrefly 测试循环分开,只跑这些验证时:
python3 test.py --no-fmt --no-lint --no-test --tensor-shapes --no-conformance --no-jsonschema运行时测试(Runtime Tests)
运行时测试验证"注解帮助函数与可运行示例模型在 Python 中行为正确",而不只是通过 Pyrefly 静态检查。测试位于:
tensor-shapes/pyrefly-torch-stubs/test/runtime_tests/运行时测试与静态 fallback 检查都需要共享 virtualenv(同一环境同时提供 torch、numpy、jax)。Bootstrap 是唯一下载依赖的步骤,因此也是唯一需要网络访问的步骤:
python3 tensor-shapes/bootstrap_venv.py # 内部机器经由 fwdproxy 追加 --fwdproxy python3 tensor-shapes/run_tests.py --runtime-only- virtualenv 默认位于
~/.tensor-shapes-venv;设置$TENSOR_SHAPES_VENV可放到别处。 - 各运行器从不创建virtualenv,也从不触网:缺失时它们会明确说明并打印 bootstrap 命令。
- Torch 与 NumPy 的静态检查使用已安装库的定义;JAX 静态检查不需要 virtualenv。
迭代时单跑某个 suite:
python tensor-shapes/pyrefly-torch-stubs/run_runtime_tests.py --suite annotation python tensor-shapes/pyrefly-torch-stubs/run_runtime_tests.py --suite model运行时运行器会为shape_extensions与可运行示例模块设置 import 路径。运行时测试在内部 checkout 中完全相同:它们针对 virtualenv 运行,从不经过 Buck,从而保证没有任何工作流会重新构建 torch、numpy 或 jax。
内核测试(Kernel Tests)
多数贡献者不需要本节。只有当改动的是 Pyrefly 的张量形状内核(而非仅 stub 或示例)时才使用这些测试。内核变更包括:
shape_extensions原语或装饰器;assert_shape的类型检查器行为;@shape_dsl_function的解析、校验或求值;@uses_shape_dsl的处理;- Pyrefly Rust 源码中的 special handlers。
聚焦的 Pyrefly 单元测试位于 pyrefly/lib/test/shape_dsl.rs,其中保留的 V1 兼容路径测试隔离在该文件的legacy模块内,并使用私有内存 stub。用 Cargo 运行:
cargo test shape_dsl内核测试刻意比 stub/示例 suite 小得多:它覆盖核心原语与不变量,而张量形状 stub 测试通过真实的 PyTorch 签名对 DSL 施压。
提交前检查(Pre-Commit Checks)
tensor-shape 包中的 Python 文件使用Ruff 格式化器而非 Black,从仓库根目录以与 CI 相同的 Ruff 版本格式化:
uv tool run --from ruff==0.16.5 ruff format \ tensor-shapesskills目录是文档而非语料源码,之所以被排除,是因为 Ruff 还会格式化 Markdown 中内嵌的 Python 片段。
在移交变更之前,还要运行仓库级格式化与 lint:
./test.py --no-test --no-tensor-shapes --no-conformance --no-jsonschema以及按所触达文件运行相应的张量形状检查:
- Stub/测试/示例变更:
python3 tensor-shapes/pyrefly-torch-stubs/run_pyrefly.py - 运行时帮助函数或可运行模型变更:
python tensor-shapes/pyrefly-torch-stubs/run_runtime_tests.py - 内核变更:
cargo test shape_dsl(或上述 Buck 等价方式)
小结:选择正确的贡献路径
| 你想做的事 | 机制 | 主要改动位置 | 验证命令 |
|---|---|---|---|
覆盖nn.Xxx、torch.xxx的形状 | Fixture stub | torch-stubs/nn/、torch-stubs/*.pyi | run_pyrefly.py --suite torch-positive |
| 表达需要计算的形状(reduction、conv、padding…) | 类型级 DSL | torch-stubs/_shapes.pyi | run_pyrefly.py(全量) |
| 用真实模型验证覆盖质量 | 移植模型 | pyrefly-torch-stubs/examples/ | verify_port.sh+--suite torch-examples |
| 校验运行时注解行为 | 运行时测试 | test/runtime_tests/ | run_runtime_tests.py |
| 改形状内核原语 | Special handlers / 内核 | pyrefly/lib/(Rust) | cargo test shape_dsl |
这条"stub 优先、DSL 次之、内核兜底"的分工正是 TENSOR_SHAPES_CONTRIBUTING.md 的核心主张:外部贡献应停留在 stub-only 或 example/test-only 层面,内核变更属于更窄的工作流。遵循这一分层,你可以不触碰 Rust 内部实现,就把 Pyrefly 的 PyTorch 形状覆盖稳步扩大。
【免费下载链接】pyreflyA fast type checker and language server for Python项目地址: https://gitcode.com/GitHub_Trending/py/pyrefly
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考