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广播直接批量计算每个维度的外积项,不需要循环取列:
- 调整
U和W的维度触发广播:将U扩展为(6, 2, 1)形状,W.T扩展为(1, 2, 5)形状,二者相乘会直接得到形状为(6, 2, 5)的数组,正好对应原循环里每个i对应的U[:,[i]] @ W.T[[i],:]外积结果 - 对这个三维数组的中间维度(即原列维度
i)应用自定义非线性函数 - 沿中间维度求和,即可得到最终形状为
(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
相关产品推荐
相关产品推荐

