PyTorch矩阵乘法结果随张量尺寸与计算图关联状态变化问题
问题:PyTorch中批量矩阵乘法与循环切片乘法结果不一致的原因
我编写了如下代码,将张量X与矩阵C相乘。对比批量乘法与循环遍历X的每个切片进行乘法的结果时,发现结果会随X的尺寸以及C是否关联计算图而不同。
import torch from torch import nn for X,C in [(torch.rand(8, 50, 32), nn.Parameter(torch.randn(32,32))), (torch.rand(16, 50, 32), nn.Parameter(torch.randn(32,32))), (torch.rand(8, 50, 32), nn.Parameter(torch.randn(32,32)).detach()) ]: # 逐个切片相乘 A = torch.empty_like(X) for t in range(X.shape[1]): A[:,t,:] = (C @ X[:,t,:].unsqueeze(-1)).squeeze(-1) # 批量相乘 A1 = (C @ X.unsqueeze(-1)).squeeze(-1) print('equal:', (A1 == A).all().item(), ', close:', torch.allclose(A1, A))
代码输出:
equal: False , close: False equal: True , close: True equal: True , close: True
我原本预期三种情况的结果都完全一致,请问这是什么原因?
参考环境信息:
import sys, platform print('OS:', platform.platform()) print('Python:', sys.version) print('Pytorch:', torch.__version__)
输出:
OS: macOS-14.4.1-arm64-arm-64bit Python: 3.12.1 | packaged by conda-forge | (main, Dec 23 2023, 08:01:35) [Clang 16.0.6 ] Pytorch: 2.2.0
原因分析
这是因为PyTorch的自动微分机制和算子优化策略导致的计算路径差异:
- 计算图追踪与算子选择:当C是
nn.Parameter(关联计算图)且X的批量维度(第一维)较小时(比如8),循环内的切片乘法会生成独立计算图节点,而批量乘法使用融合算子,两者因浮点运算顺序、硬件加速(如Apple Silicon的Neon指令集)差异产生可观测的数值偏差。 - 批量尺寸触发优化:当X第一维增大到16时,PyTorch会对循环操作自动向量化,让循环版本与批量乘法的计算路径对齐,结果完全匹配。
- detach()切断计算图:C被
detach()后,计算图被移除,PyTorch切换到纯张量运算模式,循环和批量乘法使用相同基础算子,不会因自动微分逻辑引入差异,结果一致。
另外,==要求浮点值严格匹配,torch.allclose()允许微小容差,但第一种情况连allclose都返回False,说明差异已超出正常浮点误差范围,核心是计算路径不同导致的结果偏差。
内容的提问来源于stack exchange,提问作者dkv
相关产品推荐
相关产品推荐

