三维矩阵场景下如何用numpy矩阵运算替代多层for循环实现指定公式
实现方案
你需要的核心是对三个矩阵的公共维度K做乘积求和,以下是两种无循环的numpy实现方式:
方法1:使用np.einsum(最推荐)
einsum可以直接指定多维度的求和规则,完全匹配你的原始计算逻辑,代码仅需一行:
gamma_dashed_lft = np.einsum('lk, fk, kt -> lft', q_lk, w_fk, h_kt)
参数说明:
lk、fk、kt分别对应三个输入矩阵q_lk、w_fk、h_kt的维度标识- 三个输入里重复出现的维度
k,就是需要沿该维度做乘积求和的维度 - 箭头后的
lft是输出结果的维度顺序,和你需要的(L, F, T)完全匹配
如果你的输入维度较大,可以加optimize=True参数开启底层运算优化,进一步提升速度:
gamma_dashed_lft = np.einsum('lk, fk, kt -> lft', q_lk, w_fk, h_kt, optimize=True)
方法2:使用广播+维度求和
如果不熟悉einsum的语法,也可以通过维度扩展广播相乘后求和实现:
gamma_dashed_lft = (q_lk[:, None, :, None] * w_fk[None, :, :, None] * h_kt[None, None, :, :]).sum(axis=2)
逻辑说明:
- 用
None给每个矩阵插入额外维度,让三个矩阵的维度都扩展为(L, F, K, T)后对应位置相乘 - 最后沿第2个维度(也就是K维度)求和,得到最终的
(L, F, T)结果
内容的提问来源于stack exchange,提问作者skeleton32
相关产品推荐
相关产品推荐

