ARTICLE DETAIL

资讯详情

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

详解 PyTorch `out=` 契约:输出张量的调整大小、safe copy 规则与源码实现

详解 PyTorch `out=` 契约:输出张量的调整大小、safe copy 规则与源码实现 详解 PyTorchout契约输出张量的调整大小、safe copy 规则与源码实现【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch本文基于 PyTorch 仓库中的官方开发者笔记 out.md系统讲解out参数背后完整的契约语义输出张量的 shape 不匹配时如何 resize 或报错、什么是safe copy以及它与传统 copy 的区别、type kind类型族排序如何决定 dtype 合法性并结合 c10/core/ScalarType.h 中的canCast实现验证底层检查逻辑。读完后你可以准确预测out在各种边界条件下的行为并在自定义算子或性能调优中正确、安全地使用输出张量。1.out契约总览当用户向一个算子传入一个或多个张量作为out时PyTorch 对每个输出张量遵循如下契约以下规则对多输出算子中的每一个 out 张量独立适用场景行为out 张量为空0 个元素被resize为计算结果的 shape、stride 与 memory formatout 张量 shape 与计算结果不一致抛出错误或者像空张量一样被 resize后一种 resize 行为已被弃用PyTorch 正在把各算子统一为一致的报错行为out 张量 shape 正确数值上等价于先执行运算再把结果 safe copy 进 out 张量此时 out 张量的stride 与 memory format 保持不变out 张量参与梯度requires grad不支持需要特别强调的是最后一条隐含语义即使 shape 正确PyTorch 也不改变out 张量已有的 stride 和 memory format例如 channels_last 布局的 out 张量会保持 channels_last 布局只在数值层面完成写入。2. safe copy 的定义它和普通 copy 有什么不同文档指出shape 正确的 out 张量在数值上等价于执行运算后把结果 safe copy 过去。而safe copy 与 PyTorch 的常规 copy 语义不同其核心约束来自算子是否参与类型提升type promotion不参与类型提升的算子源张量计算结果与目标张量out的device 和 dtype 必须完全一致参与类型提升的算子copy 允许落到不同 dtype 的目标上但目标不能处于比源更低的 type kind。PyTorch 定义了四个 type kind按序排列为boolean integer float complex由此可以推出一个典型例子add参与类型提升当你给两个 float 输入却传入一个integer的out张量时由于 int 的 type kind 低于 float会抛出运行时错误。反之float 计算结果写入 complex 的 out 张量是允许的目标 kind 更高。这条规则保证了 in-place / out-of-place 语义的一致性out写入本质上是一种放宽版的 in-place 赋值不允许把高精度结果安全地丢进低精度族。3. 源码印证c10::canCast就是 safe copy 的 dtype 判据上述type kind规则在 C 层由 c10/core/ScalarType.h 中的canCast(from, to)函数精确落地。该函数实现了三条硬性禁止complex → 非 complex 禁止例如float_tensor * complex这类把复数结果写进实数张量的操作被拒绝对应目标 kind 不能低于源中的 complex 最高位float → 整型禁止int_tensor * float被拒绝正是文档中add例子float 输入 integer out 报错的底层依据非 bool → bool 禁止bool被当作独立的类型类别处理以与类型提升规则保持一致——bool_tensor 5会提升到int64因此把非 bool 结果写进 bool 张量如bool_tensor 5不允许。三条检查之外的其余组合均视为可安全转换同族内的宽度变化、bool → 任意等。该函数在算子侧被实际调用。例如线性代数工具函数 aten/src/ATen/native/LinearAlgebraUtils.h 中的checkLinalgCompatibleDtype它用c10::canCast(input.scalar_type(), result.scalar_type())校验_out变体的 out 张量 dtype不满足时报错信息形如Expected result to be safely castable from ... dtype, but got result with dtype ...。这说明safe copy不只是文档约定而是许多算子在运行时真实执行的检查逻辑。4. 性能视角out 张量的存储复用与拷贝融合文档特别提醒虽然 shape 正确时out的数值语义是运算 safe copy但在实现层面许多算子如add会直接复用 out 张量的 storage 并把 copy 融合进计算省掉一次额外的内存往返。这一点对理解行为差异很重要从用户视角你只需按先算再拷贝的心智模型推理数值结果从性能视角传入正确 shape 的out张量尤其是 channels_last 布局的 pre-allocated 输出往往比拿到结果再手动 copy_更快因为底层 kernel 直接把结果写进目标存储这也解释了为什么空 out 张量会被 resize——算子倾向于让计算直接落进目标存储。5. 契约执行的现实状况与使用建议文档最后有一句重要的免责声明PyTorch 中许多算子并未正确实现上述 out 契约仍需要逐个更新。因此在使用时建议依赖 resize 行为时保持保守空张量自动 resize 是稳定的但shape 不匹配时静默 resize已被弃用代码中应始终显式传入与结果 shape 一致的 out 张量或空张量不要把正确性寄托在旧版 resize 兼容行为上out 张量不要参与 autogradrequires_gradTrue的 out 张量不受支持dtype 选择参考 safe copy 规则让 out 张量的 type kind 不低于运算结果device 保持一致拿不准时用c10::canCast的三条禁止规则对照检查见 c10/core/ScalarType.h多输出算子逐一核对如div(r1, r2, out(o1, o2))中每个 out 张量独立适用全部规则验证行为仓库测试 test/test_out_dtype_op.py 覆盖了 out 变体的 dtype 相关行为可作为自定义场景的参考基线。6. 小结out是 PyTorch 中一个看似简单、实则约束丰富的参数。其契约可以浓缩为三句话shape 正确时是运算 safe copystride 与 memory format 保持safe copy 以 bool int float complex 的类型族顺序禁止降族写入底层由c10::canCast实现而高性能实现会通过复用 out 存储、融合拷贝来超越这一朴素语义。理解这套规则既能避免 dtype/shape 边界的踩坑也能写出更省内存的数值代码。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表