You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.09 08:25:23