PyTorch原地add_操作广播形状不匹配报错原因咨询(对比Numpy)
PyTorch原地广播操作报错原因解析
你遇到的问题核心在于PyTorch的原地(in-place)操作对广播的限制和Numpy不同。
先看你的PyTorch代码:
import torch x=torch.empty(1,3,1) y=torch.empty(3,1,7) (x.add_(y)).size()
执行后报错:
RuntimeError: output with shape [1, 3, 1] doesn't match the broadcast shape [3, 3, 7]
而对应的Numpy代码可以正常得到预期的广播结果:
import numpy as np x=np.empty((1,3,1)) y=np.empty((3,1,7)) (x + y).shape # Output: (3, 3, 7)
原因解释:
PyTorch的原地操作(比如add_这类后缀带下划线的方法)要求操作后的结果形状必须和原张量的形状完全一致,不允许广播后改变原张量的形状。
虽然按照广播规则,x和y可以广播为(3,3,7)的形状进行相加,但add_是原地修改x,而x原本的形状是(1,3,1),无法容纳广播后(3,3,7)的结果,所以直接抛出形状不匹配的错误。
而Numpy的x + y是创建新的张量来存储结果,不是原地修改原数组,所以可以正常完成广播计算,生成新形状的数组。
如果要在PyTorch中实现相同的效果,应该使用非原地的add方法(或者+运算符),这样会返回一个新的张量:
import torch x=torch.empty(1,3,1) y=torch.empty(3,1,7) (x + y).size() # 或者 x.add(y).size() # 输出: torch.Size([3, 3, 7])
内容的提问来源于stack exchange,提问作者nothingissomething
相关产品推荐
相关产品推荐

