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

NumPy中自定义非线性矩阵乘法的简洁实现方法

问题描述

假设我有两个矩阵 U 和 W:

import numpy as np
U = np.arange(6*2).reshape((6,2))
W = np.arange(5*2).reshape((5,2))

标准线性矩阵乘法可以直接用如下写法实现:

U @ W.T

运行得到输出:

array([[  1,   3,   5,   7,   9],
       [  3,  13,  23,  33,  43],
       [  5,  23,  41,  59,  77],
       [  7,  33,  59,  85, 111],
       [  9,  43,  77, 111, 145],
       [ 11,  53,  95, 137, 179]])

线性场景下也可以通过按列计算、循环求和的方式得到完全一致的结果:

def mult(U, W, i):
  return U[:, [i]] @ W.T[[i],:]

sum([mult(U, W, i) for i in range(2)]) #1

输出和直接矩阵乘法完全相同:

array([[  1,   3,   5,   7,   9],
       [  3,  13,  23,  33,  43],
       [  5,  23,  41,  59,  77],
       [  7,  33,  59,  85, 111],
       [  9,  43,  77, 111, 145],
       [ 11,  53,  95, 137, 179]])

现在将mult()替换为自定义非线性函数,示例如下:

def mult(U, W, i):
  return (U[:, [i]] @ W.T[[i],:]) * np.cos(U[:, [i]] @ W.T[[i],:])

sum([mult(U, W, i) for i in range(2)]) #2

可以验证该计算结果和(U @ W.T) * np.cos(U @ W.T)的结果不一致。需要找到更简洁的写法实现#2的计算逻辑,最好能兼顾运算效率(当前处理的矩阵规模不大)。


解答

核心逻辑

你的计算本质是对每个共享维度的切片单独做非线性变换,再沿共享维度求和,和先算整体内积再做非线性变换的逻辑有本质区别,因此不能直接套用普通矩阵乘法的写法,但可以通过numpy的广播机制去掉Python层面的显式循环,写法更简洁,运算效率也更高。

实现方法

利用numpy广播直接批量计算每个维度的外积项,不需要循环取列:

  1. 调整U和W的维度触发广播:将U扩展为(6, 2, 1)形状,W.T扩展为(1, 2, 5)形状,二者相乘会直接得到形状为(6, 2, 5)的数组,正好对应原循环里每个i对应的U[:,[i]] @ W.T[[i],:]外积结果
  2. 对这个三维数组的中间维度(即原列维度i)应用自定义非线性函数
  3. 沿中间维度求和,即可得到最终形状为(6,5)的结果矩阵

对应示例代码如下:

# 批量计算每个列维度i对应的外积项,输出形状(6, 2, 5)
outer_terms = U[:, :, None] * W.T[None, :, :]
# 应用非线性变换后沿共享维度求和
result = np.sum(outer_terms * np.cos(outer_terms), axis=1)

运行后得到的结果和原循环写法完全一致。

拓展说明

  • 该写法所有运算都在numpy C底层实现,没有Python层面的循环开销,效率远高于列表推导循环写法,矩阵规模增大后效率优势会更明显
  • 如果需要替换其他自定义非线性函数,只需要修改outer_terms * np.cos(outer_terms)部分即可,比如要实现f(x) = x² + sigmoid(x),直接替换为outer_terms**2 + 1/(1+np.exp(-outer_terms))即可
  • 当共享维度数值较大时,该写法的内存开销会略高于循环写法,但你提到当前处理矩阵规模不大,完全不需要考虑该问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 21:18:26