
PyTorch torch.Size 详解张量形状的元组子类及其底层实现【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorchtorch.Size是 PyTorch 中torch.Tensor.size()的返回类型用于描述张量各维度的大小同时是 Python 内置tuple的子类因此天然支持索引、切片、长度等序列操作。本文以 size.md 为核心结合仓库中torch/csrc/Size.cpp的 C 扩展实现与torch/_C/__init__.pyi.in的类型声明系统讲解torch.Size的用法、运算语义、构造规则及其在动态符号形状SymInt场景下的行为帮助你在模型开发、测试断言与编译优化中准确使用这一类型。torch.Size 是什么torch.Size是调用torch.Tensor.size()的返回类型它描述了原始张量所有维度的大小。作为tuple的子类它支持常见的序列操作例如索引与求长度。下面的例子展示了最基本的用法与 size.md 中的示例一致 x torch.ones(10, 20, 30) s x.size() s torch.Size([10, 20, 30]) s[1] 20 len(s) 3从输出可以看出s的repr形式是torch.Size([10, 20, 30])但它本质上就是一个长度维度数为 3 的元组内部元素分别为各维大小10、20、30。与 shape 属性的关系torch.Tensor还提供了shape属性它与size()返回完全相同的对象——事实上shape就是size()的别名。两者在绝大多数场景下可以互换但size()还可以接受一个dim参数用于返回指定维度的整数大小例如x.size(1)返回20。源码级剖析torch.Size 在 C 层的实现torch.Size并非 Python 层定义的类而是在 C 扩展层实现的原生类型。其核心实现位于 torch/csrc/Size.cpp底层结构体THPSize直接内嵌了一个PyTupleObject见 Size.cpp这意味着它确实是一个“货真价实”的元组内存布局与tuple兼容类型对象THPSizeType的tp_base被设置为PyTuple_Type见 Size.cpp即显式继承自内置tuple类型这解释了为什么isinstance(s, tuple)为Truetp_repr由THPSize_repr实现见 Size.cpp它按torch.Size([a, b, c])的格式拼接各维大小这就是我们看到的repr输出的来源。因此“torch.Size是tuple的子类”不是文档中的简化说法而是类型系统层面的事实你可以在任何接受tuple的代码中直接使用torch.Size并且它还会保留自己的类型标识。构造规则与元素类型检查虽然torch.Size通常由tensor.size()返回但你也可以直接显式构造它 torch.Size([10, 20, 30]) torch.Size([10, 20, 30])其构造逻辑在THPSize_pynew中实现见 Size.cpp值得注意的行为有先按tuple的默认构造方式创建对象然后逐元素做类型规范化检查对已经是 Pythonint或可通过__index__协议转换的对象包括 0 维张量与单元素张量的元素予以接受并转换为标准整数对无法转换为整数的元素抛出TypeError错误信息为torch.Size() takes an iterable of int (item %zd is %s)额外允许SymInt符号整数元素这是为编译期动态形状场景预留的详见下文。torch.Size 支持的操作与运算语义得益于tuple基类以及 C 层的运算符重载torch.Size支持以下常用操作操作写法返回类型说明索引s[i]int/torch.SymInt取第 i 维大小切片s[1:]torch.Size切片结果仍保持torch.Size类型长度len(s)int张量维度数dim拼接s t或t storch.Size与元组拼接语义一致重复s * ntorch.Size元组重复语义元素总数s.numel()int各维大小的乘积类型声明文件 torch/_C/init.pyi.in 中完整列出了这些方法的签名class Size(tuple[_int, ...]): overload def __getitem__(self: Size, key: SupportsIndex, /) - _int: ... overload def __getitem__(self: Size, key: slice, /) - Size: ... def __add__(self, other: tuple[_int, ...], /) - Size: ... def __radd__(self: Size, other: tuple[_int, ...], /) - Size: ... def __mul__(self, other: SupportsIndex, /) - Size: ... def __rmul__(self, other: SupportsIndex, /) - Size: ... def numel(self: Size, /) - _int: ...拼接运算的细节在 C 层THPSize_add见 Size.cpp通过重载nb_add保证了tuple size也能返回torch.Size而非普通tupleTHPSize_concat则对拼接右操作数做了类型校验只允许与tuple拼接错误信息为can only concatenate tuple (not ...) to torch.Size。此外类型声明中特别注明torch.Size不支持与非整数元组相加见 torch/_C/init.pyi.in。 torch.Size([10, 20]) (30,) torch.Size([10, 20, 30]) (1, 2) torch.Size([3, 4]) torch.Size([1, 2, 3, 4])numel() 的实现numel()在 Size.cpp 中实现从 1 开始将元组中的每一维大小累乘。因此torch.Size([10, 20, 30]).numel()返回6000与torch.numel(x)及x.numel()的结果一致。在序列化与反序列化方面C 层还提供了__reduce__见 Size.cpp保证torch.Size可以被 pickle 正确还原。动态符号形状SymInt支持在torch.compile、torch.export等编译优化流程中张量的某些维度可能是符号化的torch.SymInt此时torch.Size的元素不再是普通整数。THPSize_NewFromSymSizes见 Size.cpp专门处理这种情况若元素是符号值si.is_symbolic()直接将该SymInt对象放入torch.Size从而在编译期保留形状的符号约束若元素是已知的整数值maybe_as_int()有值则展开为普通int64_t在 JIT Tracing 场景下尺寸会以 0 维张量形式记录通过torch::jit::tracer::getSizeOf因此 trace 得到的模型中形状信息可被跟踪。这也意味着在编写涉及形状处理的代码时不要假设size()返回的元素一定是 Pythonint——在动态形状编译场景下它可能是SymInt而repr实现THPSize_repr也已对SymInt元素做了分支处理见 Size.cpp以保证输出格式一致。实战建议与测试中的典型用法torch.Size在 PyTorch 代码库与测试中被广泛使用常见的实战模式包括形状断言在自定义模块的forward中校验输入形状例如判断x.size()的最后一个维度是否符合预期测试断言仓库测试中大量使用torch.Size作为期望值例如 test/ao/sparsity/test_composability.py 中的self.assertEqual(mod(torch.randn(1, 4, 4, 4)).shape, torch.Size([1, 4, 4, 4]))直接比较张量shape与字面量torch.Size这是最简洁、可读性最高的断言写法元素总数计算需要批量元素数量时优先使用x.numel()而非手动math.prod(x.size())维度解包由于torch.Size是元组可直接解包如n, c, h, w x.size()与切片配合x.size()[1:]返回torch.Size可作为view、reshape等方法的参数。小结torch.Size是 PyTorch 描述张量形状的统一类型对外它表现为支持索引、长度、拼接、重复与numel()的元组子类对内它由 C 扩展层实现继承自原生tuple并针对拼接、类型检查、序列化以及 SymInt 符号形状做了专门处理。掌握它的语义能让你在形状校验、测试断言与动态形状编译场景中写出更健壮的代码。进一步阅读size.md原始 API 文档、Size.cppC 层实现、torch/_C/init.pyi.in类型声明。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考