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

求numpy einsum对应的tensordot实现,困惑于ii->i张量缩并

用tensordot替代einsum实现'mki,tik->tim'操作

首先要澄清你提到的“缩并操作ii->i”的误解:在你的einsum表达式'mki,tik->tim'里,索引i并没有被缩并,而是作为保留维度存在的——它在两个输入张量中对应相同的维度位置,计算时会逐i进行操作,最终保留到输出里。真正被缩并求和的是索引k(只在输入中出现,输出里没有)。

核心计算逻辑

我们的目标是计算:

tim[t,i,m] = sum_k mki[m,k,i] * tik[t,i,k]

由于稀疏库不支持einsum,我们可以结合轴变换、维度合并/拆分和tensordot来实现等价操作,具体步骤如下:


方法1:基于reshape和tensordot的实现

这种方法通过合并维度把批量操作转化为tensordot能直接处理的形式:

import numpy as np

# 初始化示例数据
mki_shape = (25,25,121)
mki = np.random.uniform(size=mki_shape)
tik_shape = (10,121,25)
tik = np.random.uniform(size=tik_shape)

# 原einsum结果作为基准
tim_einsum = np.einsum('mki,tik->tim', mki, tik)

# tensordot等价实现
# 1. 调整mki的轴顺序为(m, i, k),让i维度和tik的i对齐
mki_trans = mki.transpose(0, 2, 1)
# 2. 合并m和i维度,变成(m*i, k)
mki_reshaped = mki_trans.reshape(-1, mki_trans.shape[-1])
# 3. 合并tik的t和i维度,变成(t*i, k)
tik_reshaped = tik.reshape(-1, tik.shape[-1])

# 4. 用tensordot缩并k轴(等价于矩阵乘法 tik_reshaped @ mki_reshaped.T)
tensordot_result = np.tensordot(tik_reshaped, mki_reshaped, axes=1)

# 5. 拆分维度,恢复(t, i, m)的形状
tim_tensordot = tensordot_result.reshape(tik_shape[0], tik_shape[1], mki_shape[0])

# 验证结果一致性
print(np.allclose(tim_einsum, tim_tensordot))  # 输出True

方法2:基于批量矩阵乘法的实现(推荐,如果稀疏库支持)

如果你的稀疏张量库支持批量矩阵乘法(比如sparse.matmul),可以用更直观的方式实现,不需要reshape:

# 调整轴顺序,把i作为批量维度
mki_batch = mki.transpose(2, 0, 1)  # 形状(121, 25, 25),对应(i, m, k)
tik_batch = tik.transpose(1, 0, 2)  # 形状(121, 10, 25),对应(i, t, k)

# 对每个i执行批量矩阵乘法:tik_batch[i] @ mki_batch[i].T
tim_batch = np.matmul(tik_batch, mki_batch.transpose(0, 2, 1))

# 调整轴顺序回到(t, i, m)
tim_matmul = tim_batch.transpose(1, 0, 2)

# 验证一致性
print(np.allclose(tim_einsum, tim_matmul))  # 输出True

这两种方法都能得到和einsum完全一致的结果,且适配稀疏张量的操作需求。

内容的提问来源于stack exchange,提问作者Tom

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.07 18:32:30