ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

PyTorch 测试工具库 torch.testing 完全指南:assert_close、make_tensor 与 assert_allclose 实战解析

2026/9/10 10:29:11 拓冰建站 浏览量
PyTorch 测试工具库 torch.testing 完全指南:assert_close、make_tensor 与 assert_allclose 实战解析 PyTorch 测试工具库 torch.testing 完全指南assert_close、make_tensor 与 assert_allclose 实战解析【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorchtorch.testing 是 PyTorch 官方提供的测试工具子模块它把如何严谨地比较张量与构造测试张量这两个测试中最基础、也最容易出错的问题封装成了开箱即用的标准 API。本文围绕官方文档 docs/source/testing.md 中公开的三个核心接口展开结合仓库内 torch/testing/_comparison.py 与 torch/testing/_creation.py 的源码实现完整讲解默认容差表、比较判定公式、全部参数语义与真实错误信息样例帮助你写出比裸用 / numpy.allclose更严谨、更可维护的 PyTorch 测试代码。torch.testing 模块概览torch.testing从模块层级上属于 PyTorch 测试基础设施的一部分它的公共 API 在 torch/testing/init.py 中统一导出from ._comparison import assert_allclose, assert_close as assert_close from ._creation import make_tensor as make_tensor也就是说用户可见的完整接口为三个函数接口职责状态torch.testing.assert_close按容差严格比较两个张量/标量/序列/映射推荐使用torch.testing.make_tensor按指定 shape、dtype、device 构造均匀分布的随机测试张量推荐使用torch.testing.assert_allclose旧版近似比较接口自 1.12 起弃用从源码结构看比较逻辑全部收敛在_comparison.py的Pair抽象类体系与originate_pairs递归展开机制中而张量构造收敛在_creation.py的make_tensor一个函数内。理解这两个文件的内部约定是吃透该模块的关键。assert_close严谨的数值容差比较torch.testing.assert_close是当前 PyTorch 官方推荐的首选比较接口它在 torch/testing/_comparison.py 中实现。其设计目标非常明确默认严格参数可配错误信息自带诊断统计。判定公式与容差表对于 strided、非量化、实值且有限的张量assert_close认为两者接近的数学定义是|actual - expected| atol rtol * |expected|非有限值-inf与inf只有在彼此相等时才被认为接近NaN只有在equal_nanTrue时才被认为彼此相等。当用户未显式传入rtol/atol时容差由输入张量的dtype决定官方文档即该函数 docstring 中的表格给出的默认值如下dtypertolatoltorch.float161e-31e-5torch.bfloat161.6e-21e-5torch.float321.3e-61e-5torch.float641e-71e-7torch.complex321e-31e-5torch.complex641.3e-61e-5torch.complex1281e-71e-7torch.quint81.3e-61e-5torch.quint2x41.3e-61e-5torch.quint4x21.3e-61e-5torch.qint81.3e-61e-5torch.qint321.3e-61e-5其他 dtype0.00.0当actual与expected的 dtype 不一致时取两者容差中更宽松更大的一组。这一表格并非文档凭空撰写其数据源就在 torch/testing/_comparison.py 顶部的_DTYPE_PRECISIONS字典中量化 dtype 的容差被显式对齐到 float32因为量化张量会先dequantize再以浮点表示比较。取最宽松容差的逻辑则由default_tolerancesL115-L140用max(rtols), max(atols)实现。一个值得注意的约束rtol与atol必须同时指定或同时省略否则会抛出ValueError。这样设计的原因在源码注释中写得很直白——只指定atol0.0时张量仍可能因为rtol0而意外通过。参数全解assert_close的完整签名以当前仓库源码为准assert_close( actual, expected, *, allow_subclassesTrue, rtolNone, atolNone, equal_nanFalse, check_deviceTrue, check_dtypeTrue, check_layoutTrue, check_strideFalse, msgNone, )各参数语义如下actual/expected任意输入。可以是torch.Tensor也可以是能用torch.as_tensor构造出张量的任何 tensor-like 或 scalar-like 对象含 numpy 数组与 Python 标量还可以是collections.abc.Sequence或Mapping此时会按结构递归逐元素比较。除 Python 标量外两个输入的类型必须直接相关一方是另一方的实例反之亦然。allow_subclasses默认True是否允许直接相关但类型不同的输入如torch.nn.Parameter与torch.Tensor。设为False时要求类型完全一致。rtol/atol默认None相对/绝对容差必须同时给出省略时使用上表默认值。equal_nan默认False为True时两个NaN视为相等。check_device默认True是否校验对应张量在相同设备上关闭后不同设备上的张量会被搬到 CPU 再比较。check_dtype默认True是否校验 dtype 一致关闭后不同 dtype 会按torch.promote_types提升到公共 dtype 再比较。check_layout默认True是否校验 layout 一致关闭后不同 layout 会先转成 strided 张量再比较。check_stride默认False为True且张量为 strided 时额外校验 stride 是否一致。msg默认None自定义失败信息也可以传入可调用对象它会收到自动生成的信息并返回新信息用于在错误前后追加上下文。支持的张量形态assert_close并非只能比较普通稠密张量TensorLikePair._compare_values 展示了其完整的分支策略量化张量要求is_quantized状态与qscheme一致然后对dequantize()后的结果做接近性比较稀疏 COO依次比较稀疏维数、nnz、indices相等性与values接近性稀疏压缩格式CSR / CSC / BSR / BSC比较nnz、compressed/plain indices相等性并先统一索引 dtype与values接近性JaggedNJT直接对values()做常规接近比较float8 等 1 字节浮点强制rtol0、atol0按位比较meta 张量只做属性检查不触碰数据。错误信息自带诊断失败时assert_close抛出的AssertionError会包含不匹配元素比例 最大绝对/相对差 首处不匹配索引例如官方示例Tensor-likes are not close! Mismatched elements: 2 / 3 (66.7%) Greatest absolute difference: 2.0 at index (1,) (up to 1e-05 allowed) Greatest relative difference: 1.0 at index (1,) (up to 1.3e-06 allowed)这条信息的生成路径在 make_tensor_mismatch_msg内部先计算torch.isclose的布尔掩码再把不匹配元素置零后求最大差值从而保证统计量只反映真正的失配点。若需要为团队自定义统一的断言风格官方建议用functools.partial固化默认参数例如import functools assert_equal functools.partial(torch.testing.assert_close, rtol0, atol0)比较器的内部架构Pair 体系assert_close底层并没有手写一堆 if-else而是由 Pair 抽象基类派生出四个具体比较器NonePair处理NoneBooleanPair处理bool及 numpy 的np.bool_NumberPair处理int/float/complex及np.number并按int→int64、float→float64、complex→complex128的映射选取默认容差TensorLikePair处理所有 tensor-like 输入承担属性检查、属性对齐与各类特殊 layout 的值比较ObjectPair兜底比较器使用运算符只应作为最后手段。originate_pairsL1158-L1340负责把嵌套的 Sequence / Mapping / dataclass 递归展开成扁平的一串Pair并沿途携带id元组如(0, layer, 2)这样出错时可以精确定位到容器中的哪一个元素。assert_close只抛出第一个失败 pair 的ErrorMeta源码中not_close_error_metasL1343-L1412会收集所有失败项并在返回前主动打破ErrorMeta的引用环以避免测试中的 CUDA 显存泄漏——这是从实际测试工程中沉淀出来的细节。make_tensor一行代码构造可复现的测试张量torch.testing.make_tensor用于在指定设备、dtype 下快速生成值均匀分布于[low, high)区间的张量实现在 torch/testing/_creation.py是 PyTorch 各测试套件如test/目录下大量用例构造随机输入的标准工具。签名与默认取值范围make_tensor( *shape, dtype, device, lowNone, highNone, requires_gradFalse, noncontiguousFalse, exclude_zeroFalse, memory_formatNone, )当low/high缺省时取值区间由dtype决定源码 docstring 中的官方表格dtypelowhigh布尔类型02无符号整型010有符号整型-910浮点类型-99复数类型-99shape既支持make_tensor(3, 3, ...)这样的可变参数形式也支持make_tensor((3, 3), ...)的单序列形式源码在 L178-L180 做了归一化。dtype与device是必填参数。参数细节与边界行为low/high显式传入时若超出该 dtype 可表示的有限范围会被自动钳制到最值传入nan会抛ValueErrorlow high也会抛ValueError。源码用modify_low_highL127-L176统一完成钳制并对整型采用ceil保证采样值不越界。requires_grad默认False浮点与复数类型支持布尔/整型上设为True会抛ValueError。noncontiguous默认False返回非连续张量。实现技巧是先放大最后一维再步进切片L193-L197 与 L262-L264目的是顺带覆盖 offset 相关的边界问题元素少于 2 个时该参数被忽略且与memory_format互斥。exclude_zero默认False把采样到的 0 替换掉——布尔/整型替换为 1浮点替换为该 dtype 的torch.finfo(dtype).tiny最小正规格化数复数替换为实部虚部均为tiny的复数。该参数对除法、log、归一化类算子的测试尤为实用。memory_format指定返回张量的内存格式与noncontiguous互斥。实现要点从源码看make_tensor内部按 dtype 家族分流布尔/整型走torch.randint浮点与复数先torch.empty再_uniform_random_填充复数通过torch.view_as_real对视作实张量采样float8 家族先在 float32 上采样再to(dtype)转换。_uniform_random_L36-L42) 还有一个值得学习的细节当high - low超过 dtype 最大值时先采样一半区间再整体乘 2规避底层uniform_对区间宽度的限制。典型用法from torch.testing import make_tensor # 值域 [-1, 1) 的 float32 张量 x make_tensor((3,), devicecpu, dtypetorch.float32, low-1, high1) # CUDA 上的 bool 张量 mask make_tensor((2, 2), devicecuda, dtypetorch.bool) # 参与自动求导的浮点张量 w make_tensor((4, 4), devicecpu, dtypetorch.float64, requires_gradTrue) # 构造不含 0 的张量适合除法/对数类算子测试 y make_tensor((8,), devicecpu, dtypetorch.float32, exclude_zeroTrue)assert_allclose已弃用的旧接口torch.testing.assert_allclose自 PyTorch 1.12 起被标记为弃用FutureWarning并计划在未来的版本中移除完整实现在 torch/testing/_comparison.py。它存在两个与assert_close显著不同的默认行为equal_nan默认Trueassert_close为False默认容差按旧约定取值float16 为(1e-3, 1e-3)、float32 为(1e-4, 1e-5)、float64 为(1e-5, 1e-8)且比较时强制check_deviceTrue、check_dtypeFalse、check_strideFalse。源码实际上只是把参数整理后转调assert_close因此迁移成本极低。官方给出的升级指引是直接改用assert_close并在迁移时注意两处语义变化一是显式检查 dtype 是否一致必要时传check_dtypeFalse保持旧行为二是按需传equal_nanTrue同时建议使用assert_close按 dtype 自动选择的默认容差而不是沿用旧值。如果你的测试库中仍在使用该函数应尽快替换以消除弃用警告。在真实测试中的落地建议综合文档与源码可以提炼出几条可操作的最佳实践用assert_close取代torch.allclose/numpy.testing.assert_allclose它同时覆盖张量、标量、numpy 数组、嵌套序列/映射与 dataclass错误信息自带失配统计且默认值经过 PyTorch 官方按 dtype 校准源码_DTYPE_PRECISIONS即其依据。用functools.partial固化项目级断言例如assert_equal functools.partial(torch.testing.assert_close, rtol0, atol0)用于需要严格相等的场景需要比默认更宽松的基准测试时可按需覆盖rtol/atol。用make_tensor统一测试输入构造必填device参数天然适合跨 CPU/CUDA 的参数化测试exclude_zeroTrue规避除零与对数奇异点noncontiguousTrue可以低成本覆盖内存布局类回归。优先在torch/testing/_internal/common_utils.py等测试基础设施之上编写用例整个 PyTorch 仓库的测试套件test/目录都建立在torch.testing之上遵循同样的比较与构造约定可以保证测试语义与官方一致。小结torch.testing虽然只是 PyTorch 测试栈中一个看似不起眼的模块但assert_close的容差表、make_tensor的默认值区间、以及两者围绕 dtype 与 layout 的边界处理都是官方在长期测试实践中沉淀出的标准答案。无论是编写单测、做算子数值校验还是跑模型精度回归把torch.testing用熟都能让测试代码更短、更严谨、更可诊断。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考