【发布时间】:2020-12-16 20:46:19
【问题描述】:
假设我有张量X 和Y,它们都是(batch_size, d) 维度的。我想找到由[X[0]@Y[0].T, X[1]@Y[1].T, ...] 产生的(batch_size x 1) 张量
我可以想到两种方法,但都不是特别有效。
方式 1
product = torch.eye(batch_size) * X@Y.T
product = torch.sum(product, dim=1)
这可行,但对于大型矩阵,有很多浪费的计算
方式 2
product = torch.cat(
[ X[i]@Y[i].T for i in X.size(0) ],
dim=0
)
这很好,因为没有浪费任何周期,但它不会利用任何内置的并行性 torch 提供。
我知道 numpy 有一种方法可以做到这一点,但是将张量转换为 np 数组会破坏反向传播链,这是针对神经网络的,所以这不是一个选择。
我是否缺少明显的内置 torch 方法,还是我坚持使用这两个选项?
【问题讨论】: