十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyRefly Torch Stubs 指南:为 PyTorch 提供带张量形状信息的类型桩

PyRefly Torch Stubs 指南:为 PyTorch 提供带张量形状信息的类型桩 PyRefly Torch Stubs 指南为 PyTorch 提供带张量形状信息的类型桩【免费下载链接】pyreflyA fast type checker and language server for Python项目地址: https://gitcode.com/GitHub_Trending/py/pyreflyPyRefly 的pyrefly-torch-stubs是一个以形状推断shape inference为核心的 PyTorch 类型桩包。它以 PEP 561 stub-only 发行版的形式安装torch-stubs桩包让 Pyrefly 能够在不对 PyTorch 运行时做任何替换或遮蔽的前提下发现并应用带形状信息的桩进而在编译期静态验证张量操作的形状正确性。阅读完本文你将掌握该包的设计动机、包结构与安装方式、静态检查与运行时测试的运行方法以及形状如何以类型级 DSL 的形式参与推断与报错。包是什么PEP 561 stub-only 发行版该包本质是一套「只含桩stub-only」的发行版遵循 PEP 561 约定。它安装torch-stubs桩包使 Pyrefly 能够发现针对运行时torch包的形状感知桩而不会替换或遮蔽 PyTorch 本身——运行时导入的仍是真正的torch桩只服务于静态类型检查。这一设计在源码中有两处直接体现桩包根目录 torch-stubs/py.typed 内容仅为一行partial即声明本桩只覆盖torch的一部分打包配置 pyproject.toml 中wheel 目标only-include [torch-stubs]sdist 目标include [torch-stubs/, README.md, LICENSE]确保发布物只含桩文件。包的版本与 Pyrefly 保持同步lockstep并依赖相匹配的pyrefly-shape-extensions包dependencies [ pyrefly-shape-extensions0.0.0, ]同时声明requires-python 3.12构建后端为 hatchling。partial 桩的查找语义逐模块回退partial标记带来的行为值得展开。在 torch-stubs/init.pyi 的注释中写明Generator在该包内并未定义解析它依赖「部分桩包如何被查找」的机制——当类型检查器遇到本包未定义的子模块时会回退到真实的 torch 桩typeshed 中 torch 的部分torch._C就是其中之一因此Generator来自 torch 自身。关键点是回退按模块而非按名字进行本包定义过的模块会整体遮蔽 torch 的对应版本因此__init__.pyi中保留了模块级__getattr__让尚未覆盖到的名字保持渐进式gradual即退化为Any行为。包内结构速览tensor-shapes/pyrefly-torch-stubs/ ├── torch-stubs/ # 桩本体含 nn/、distributions/、linalg/、fft/ 等子模块 │ ├── __init__.pyi │ ├── _shapes.pyi # 类型级形状 DSL 函数库 │ ├── _tensor.pyi │ ├── nn/ # nn.Module、functional、parameter、attention 等 │ ├── linalg.pyi │ ├── fft.pyi │ └── py.typed # 内容为 partial ├── examples/ # 带形状注解的模型示例resnet、bert、nanogpt 等 30 个 │ └── runtime/ # 可运行的示例含 _runnable 后缀 ├── test/ # 正向、负向、jaxtyping 与运行时测试 ├── run_pyrefly.py # 静态检查入口 ├── run_runtime_tests.py # 运行时测试入口 ├── suites.py # 定义测试套件 ├── stub_coverage.toml # 覆盖率配置 ├── pyproject.toml # 打包配置 └── pyrefly.toml # Pyrefly 搜索路径配置运行静态检查run_pyrefly.py对「本包有意未做形状标注的模块」静态检查使用已安装的 Torch 包作为回退python3 tensor-shapes/pyrefly-torch-stubs/run_pyrefly.py该运行器默认使用~/.tensor-shapes-venv虚拟环境可通过环境变量$TENSOR_SHAPES_VENV指定其他虚拟环境或使用--python显式传入一个装有 Torch 的虚拟环境解释器。解释器的解析顺序见 shape_testing.py 中venv_python显式--python→$TENSOR_SHAPES_VENV→ 默认位置~/.tensor-shapes-venv。命令行参数一览run_pyrefly.py 支持的参数如下参数作用--pyrefly PATH直接使用指定二进制这是唯一不先构建 Pyrefly 的模式--buck改用 Buck 构建并运行 Pyrefly而非 Cargo--release用 Cargo release profile 构建默认 debug--python PATH提供 Torch 回退模块的解释器默认共享虚拟环境--suite NAME只运行指定套件可重复传入默认运行全部--nocapture直接流式输出 Pyrefly 结果默认仅失败时打印Pyrefly 二进制解析逻辑Pyrefly 二进制如何被找到同样有清晰的优先级shape_testing.pypyrefly_command显式--pyrefly→--buckbuck2 run fbcode//pyrefly:pyrefly --→ 环境变量$PYREFLY→ Cargo 构建cargo build -p pyrefly [--release]。注释中特别警告复用target/里恰好存在的旧二进制是开发者调试出「看起来像真实差异的假症状」的常见原因因此除显式传二进制外其余模式都会先构建再检查。若 Cargo 不在 PATH 上会提示传入--buck或--pyrefly/$PYREFLY。测试套件定义suites.py 定义了五个套件套件名匹配文件说明torch-examplesexamples/*.py、examples/runtime/*.py带形状注解的真实模型torch-positivetest/test_*.py正向应通过检查torch-negativetest/negative_tests/test_*.py负向按# E:期望匹配错误启用expectationsjaxtyping-positivetest/jaxtyping/test_*.py使用 Python 3.12、独立 pyrefly.toml 与 fixturesjaxtyping-negativetest/jaxtyping/negative_tests/test_*.pyjaxtyping 负向套件启用expectations一个值得注意的实现细节run_pyrefly.py中调用check_suites(..., check_stubsFalse)其 TODO 注释说明——Torch 桩自身尚未启用自检self-checking因为还存在内部导入不完整、类型参数遮蔽、与注解不兼容的旧默认值等问题因此目前只对消费者套件做检查。这是「部分桩」策略的又一体现桩只对使用方负责自身仍需渐进完善。运行运行时测试run_runtime_tests.py与 numpy、jax 套件不同torch 的运行时测试是独立的 unittest 模块位于test/runtime_tests由 run_runtime_tests.py 驱动python3 tensor-shapes/pyrefly-torch-stubs/run_runtime_tests.py支持--torch-root默认包根目录与--suiteall、annotation、torchscript、model参数。三个套件分别为注解运行时行为test_annotation_runtime*.py、TorchScript 剥离器兼容模式test_torchscript_stripper_runtime.py、模型运行时test_model_runtime.py。其中torchscript被有意排除在annotation套件之外——导入shape_extensions.torchscript会在进程范围内启用兼容模式改变Int[...]行为从而影响注解测试的断言。运行时测试会额外把pyrefly-shape-extensions与examples/runtime加入sys.path。静态检查与运行时测试的双轨验证tensor-shapes 系列的测试框架采用「双重验证」设计shape_testing.py 模块 docstring每个桩包既让 Pyrefly 静态检查测试文件又让同一批文件针对真实库执行。该框架不接触网络虚拟环境仅由 bootstrap_venv.py 创建无网络权限的调用者会得到一条可操作的报错提示包括在 Meta 环境需要fwdproxy的说明。Suite数据类中的expectations标志值得关注启用后 Pyrefly 以--expectations运行把报告的错误与# E:注释匹配将「静态拒绝」与「库本身抛出的运行时错误」配对验证。由于它还会统计被抑制的错误torch 语料仅对专门的负向测试目录启用。桩的核心Tensor 与形状类型参数桩的核心是 torch-stubs/init.pyi 中的Tensor类形状以类型参数跟踪class Tensor[Shape: _Shape _Shape]:其中type _Shape IntTuple。形状推断通过注解中的类型级函数表达库特定的形状函数定义在torch/_shapes.pyi中。__init__.pyi开头的 docstring 点明「大多数形状变换由注册在类型检查器中的 meta-shape 函数处理而非在此显式给出类型签名」——这意味着常见形状操作的精确实现在 Pyrefly 类型检查器内部。形状如何参与推断以 _shapes.pyi 为例torch-stubs/_shapes.pyi 汇集了 30 个type_shape_dsl_function装饰的类型级形状函数。以reshape_shape为例它完整实现了 PyTorch reshape 的约束只允许一个-1维度、禁止小于-1的维度值、当元素总数不匹配时报错「reshape target element count does not match the input」并通过dsl.prod在已知与未知维度间计算。reduce_shape则实现了规约运算的维度归一化、越界与重复维度检查以及keepdim语义type_shape_dsl_function def reduce_shape(shape, dim, keepdim): ... if any(normalized.count(item) 1 for item in normalized): return dsl.Invalid(duplicate dimension) if keepdim: return dsl.IntTuple( (1 if index in normalized else shape[index] for index in range(len(shape))) ) return dsl.IntTuple( (shape[index] for index in range(len(shape)) if index not in normalized) )dsl.Invalid(...)携带人类可读的错误消息最终呈现为类型检查器的形状错误。其他如transpose_shape、permute_shape、squeeze_shape、expand_shape、flatten_shape、movedim_*、split_*、chunk_shapes等都以同样的模式校验参数并计算输出形状——这正是「负向测试」中错误消息的源头。Tensor 方法中的形状表达Tensor的方法签名大量使用类型级函数与Flag/IntVar/MapIntTuples等构造均来自shape_extensions。几个代表性示例def __matmul__Left: IntTuple, Right: IntTuple - Tensor[matmul_shape(Left, Right)]: ... overload def __add__OtherShape: _Shape - Tensor[broadcast(Shape, OtherShape)]: ... overload def __add__(self, other: float | int) - Self: ... overload def split... - MapIntTuples[lambda S: Tensor[S], split_size_shapes(Shape, SplitSize, Dim)]: ... def item(self: Tensor[[]]) - builtins.float | builtins.int: ...要点张量间算术返回Tensor[broadcast(...)]而非Self因为广播会改变形状特化任意子类未必保留split/chunk用MapIntTuples把形状元组映射为多个张量类型item()只接受Tensor[[]]零维才返回标量__getitem__通过index_shape(Shape, I)在类型层面计算索引后的形状。而Tensor.__getattr__返回Any表明该覆盖层只建模形状相关成员其余大规模 API 保持渐进式直到获得精确签名。nn 模块与真实模型示例torch-stubs/nn 覆盖nn.Module、nn.functional、nn.parameter、nn.init与nn.attention含 flex_attention。examples/目录下 30 个带形状注解的真实模型展示了落地效果以 examples/resnet.py 为例class ResNetBlockC: IntVar: Shape-preserving residual block: (B, C, H, W) - (B, C, H, W). def __init__(self, c: Int[C], act_fn: ShapePreservingActivation) - None: self.net nn.Sequential( nn.Conv2d(c, c, kernel_size3, padding1, biasFalse), nn.BatchNorm2d(c), act_fn(), nn.Conv2d(c, c, kernel_size3, padding1, biasFalse), nn.BatchNorm2d(c), ) ... def forwardB: IntVar, H: IntVar, W: IntVar - Tensor[[B, C, H, W]]: z self.net(x) assert_type(z, Tensor[[B, C, H, W]]) out z x return out模型用IntVar/Int[C]把批量、通道、高宽建模为符号维度forward签名精确写出Tensor[[B, C, H, W]] → Tensor[[B, C, H, W]]并用assert_type在检查时验证中间结果形状。examples/runtime/下的*_runnable.py则是可实际运行的版本同时接受静态检查与运行时验证双轨检验。覆盖率配置与限制说明stub_coverage.toml 记录了桩的覆盖率目标runtime_package torch、stub_directory torch-stubs并跳过torch._shapes该模块为类型级 DSL 函数不做成员覆盖统计成员目标覆盖torch.Tensor与torch.device。这印证了包内文档所述桩的主体工作聚焦 Tensor 的形状 API 与 device 类型其余部分以__getattr__保持渐进。当前仓库中torch 桩仍有明确的渐进边界模块级__getattr__兜底未覆盖名字、_shapes.pyi顶部有关于IntTuple切片保留符号秩用例的 TODO、split对list[int]参数暂不做形状推断列表可变V2 类型系统无法在类型层保留元素值、check_stubsFalse表明桩自检尚未启用。这些都是阅读与使用该包时值得留意的已知范围。小结pyrefly-torch-stubs把 PyTorch 的形状约束搬进了类型系统Tensor[Shape]在类型参数中携带形状_shapes.pyi以类型级 DSL 函数实现形状计算与校验examples/中的真实模型证明其可落地而partial桩机制保证它不干扰运行时torch包。通过run_pyrefly.py做静态检查、run_runtime_tests.py做运行时验证开发者可以在编译期捕获形状错误同时保持与真实 PyTorch 生态的完全兼容。进一步阅读共享测试框架 tensor-shapes/shape_testing.py、形状扩展语言 tensor-shapes/pyrefly-shape-extensions、tensor-shapes 系列总览 tensor-shapes/README.md 与 TENSOR_SHAPES_CONTRIBUTING.md。【免费下载链接】pyreflyA fast type checker and language server for Python项目地址: https://gitcode.com/GitHub_Trending/py/pyrefly创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表