NumPy einsum能否结合显式模式与编程接口实现任意维度数组指定轴收缩
NumPy的einsum整数列表编程接口本身就支持和显式模式完全一致的输出轴指定能力,不需要额外做格式转换,完全可以实现你要的任意维度数组指定轴收缩需求。
两种调用范式的逻辑是完全对应的:
- 字符串模式里
->后的字符代表最终要保留的轴标签 - 整数列表模式里,只需要在所有输入数组、对应轴映射列表的参数末尾,额外传入一个整数列表指定要保留的轴编号,就能实现和
->完全相同的显式模式效果
你之前写的np.einsum(a, [0, 1], a, [0, 1])没有传末尾的输出轴列表,默认走隐式收缩逻辑,把所有出现过的轴全部求和收缩,最终返回标量204,和省略->的字符串写法'ij,ij'效果完全一致。如果要实现和'ij,ij->j'相同的、保留1轴收缩0轴的效果,用整数接口写法如下,运行结果和字符串模式完全匹配:
>>> import numpy as np >>> a = np.arange(9).reshape((3, 3)) >>> np.einsum(a, [0, 1], a, [0, 1], [1]) array([45, 66, 93])
基于这个能力写通用收缩函数非常简单,核心逻辑就是给需要配对收缩的轴分配相同的整数标签,给不需要收缩、需要保留的轴分配互不重复的独立标签,最后把要保留的标签组成列表作为最后一个参数传入即可,参考实现:
def contract_axes(a, b, a_axis_map, b_axis_map, keep_axis_ids): """ 两个任意维度数组的指定轴收缩函数 参数: a, b: 待运算的两个NumPy数组 a_axis_map: 长度等于a维度数的整数列表,标记a每个轴的ID,两个数组中ID相同的轴会被配对收缩 b_axis_map: 长度等于b维度数的整数列表,标记b每个轴的ID keep_axis_ids: 最终输出结果中需要保留的轴ID列表 """ return np.einsum(a, a_axis_map, b, b_axis_map, keep_axis_ids)
举个实际使用的例子:如果a是形状为(2,3,4)的三维数组,b是形状为(3,4,5)的三维数组,需要收缩两个数组中长度为3、4的匹配轴(也就是a的第1、2轴和b的第0、1轴),最终保留a的第0轴、b的第2轴,输出形状为(2,5),调用方式如下:
>>> a = np.random.randn(2, 3, 4) >>> b = np.random.randn(3, 4, 5) # 待收缩的两个配对轴统一标记为1、2,待保留的a第0轴标记为0,待保留的b第2轴标记为3 >>> result = contract_axes(a, b, [0, 1, 2], [1, 2, 3], [0, 3]) >>> result.shape (2, 5)
补充一点:整数接口的轴标签不需要是从0开始的连续值,只要相同ID代表需要配对的轴、不同ID代表独立轴就可以正常运行,哪怕用负数当标签也没问题,处理高维数组的时候比拼接字符串标签更灵活,也不容易出现标签写错的问题。
内容的提问来源于stack exchange,提问作者asmeurer
相关产品推荐
相关产品推荐

