如何用torch.stack()和torch.tensordot()复现张量运算并确保结果一致?
用torch.stack()和张量运算复现逐元素乘加操作
问题描述
我希望结合torch.stack()和torch.tensordot()复现一个张量运算,以便在大型程序中进行泛化。目前已通过逐元素乘法与加法得到张量V_1,现尝试用堆叠和张量点积操作得到等价的V_2,但torch.tensordot的dims参数设置不符合预期,导致V_2与V_1形状不同、结果不相等,期望torch.all(V_1.eq(V_2))返回tensor(True)。
原实现代码
import torch N, t , J = 4, 2 , 3 K_f , K_r = 1, 1 R = 5 K = K_f + K_r id = torch.arange(N).repeat(t).sort() X = torch.randn(N*t, K , J) Y = torch.randn(N*t, 1) D = torch.randn(N, K_r , R) Draw = D.repeat_interleave(t,0) beta = torch.randn(2*K_r + K_f, 1) beta_R = (beta[0:K_r,0] + beta[K_r:2*K_r,0] * Draw ).repeat(1,J,1) print("shape beta_R:", beta_R.shape) beta_F = beta[2*K_r:2*K_r + K_f,0].repeat(N*t, J, R) print("shape beta_F:", beta_F.shape) XX_0 =X[:,0,:].unsqueeze(2).repeat(1,1,R) print("shape XX_0:", XX_0.shape) XX_1 =X[:,1,:].unsqueeze(2).repeat(1,1,R) print("shape XX_1:", XX_1.shape) V_1 = XX_0 * beta_R + XX_1 * beta_F print("shape V_1:",V_1.shape) # 输出: # shape beta_R: torch.Size([8, 3, 5]) # shape beta_F: torch.Size([8, 3, 5]) # shape XX_0: torch.Size([8, 3, 5]) # shape XX_1: torch.Size([8, 3, 5]) # shape V_1: torch.Size([8, 3, 5])
尝试复现的代码(存在问题)
# 用堆叠和tensordot复现 stack_XX = torch.stack((XX_0, XX_1), 0) print("shape stack_XX:",stack_XX.shape) stack_beta = torch.stack((beta_R, beta_F), 0) print("shape stack_beta:", stack_beta.shape) # 尝试在第一维度做张量点积 V_2 = torch.tensordot(stack_XX, stack_beta, dims=([0], [0])) print("shape V_2:",V_2.shape) # 检查是否相等 print(torch.all(V_1.eq(V_2))) # 输出: # shape stack_XX: torch.Size([2, 8, 3, 5]) # shape stack_beta: torch.Size([2, 8, 3, 5]) # shape V_2: torch.Size([8, 3, 5, 8, 3, 5]) # tensor(False)
问题原因
当前tensordot的dims参数仅指定在第0维度做张量点积,这会触发两个张量在第0维度的全量组合运算,最终得到笛卡尔积形状的结果,完全偏离了原运算的逻辑——原运算为对应位置的元素相乘后,在堆叠维度求和。
正确实现方式
方法1:逐元素相乘+堆叠维度求和(最直观高效)
这是最直接的等价实现,完全匹配原运算逻辑:
stack_XX = torch.stack((XX_0, XX_1), 0) stack_beta = torch.stack((beta_R, beta_F), 0) # 逐元素相乘后,在堆叠的第0维度求和 V_2 = torch.sum(stack_XX * stack_beta, dim=0) print("shape V_2:", V_2.shape) # 输出 torch.Size([8, 3, 5]) print(torch.all(V_1.eq(V_2))) # 输出 tensor(True)
方法2:使用tensordot实现等价运算(若必须用tensordot)
如果业务场景要求必须使用tensordot,可以将堆叠维度调整至最后,再与全1向量做张量点积(等价于求和):
# 将堆叠维度移至最后一维 stack_XX = torch.stack((XX_0, XX_1), dim=-1) stack_beta = torch.stack((beta_R, beta_F), dim=-1) # 对最后一维(堆叠维度)做张量点积,等价于求和 V_2 = torch.tensordot(stack_XX * stack_beta, torch.ones(2, device=X.device), dims=([-1], [0])) print("shape V_2:", V_2.shape) # 输出 torch.Size([8, 3, 5]) print(torch.all(V_1.eq(V_2))) # 输出 tensor(True)
内容的提问来源于stack exchange,提问作者Álvaro A. Gutiérrez-Vargas
相关产品推荐
相关产品推荐

