如何使用Numpy实现按分桶累加张量的无循环优化方案
Numpy分桶求和高效实现方案
我们有多种完全向量化的实现方案,均可以避免Python层循环,性能远高于朴素for循环实现:
方案1:基于numpy.bincount的实现
该方案利用np.bincount原生支持按索引累加权重的特性,是CPU场景下性能最优的方案之一:
- 核心逻辑:通过维度变换将高维输入适配为
bincount的输入格式,计算完成后恢复目标维度。
import numpy as np def bucket_sum_bincount(arr, idx, M): # arr形状: (B, N1, N2, ..., Nk) # idx形状: (B,) 取值范围0~M-1 # 返回形状: (N1, N2, ..., Nk, M) # 把批次B维度移到最后 arr_trans = np.moveaxis(arr, 0, -1) # 拉平非批次维度 arr_flat = arr_trans.reshape(-1, arr.shape[0]) # 按索引分桶求和,minlength保证返回长度固定为M res_flat = np.apply_along_axis( lambda x: np.bincount(idx, weights=x, minlength=M), axis=1, arr=arr_flat ) # 恢复目标维度 return res_flat.reshape((*arr_trans.shape[:-1], M))
方案2:基于np.add.at的实现
该方案代码可读性最高,逻辑灵活易修改,适合需要自定义累加规则的场景:
- 核心逻辑:初始化全零结果数组,利用Numpy高级索引和
np.add.at的重复索引原地累加特性,直接完成分桶求和。
def bucket_sum_addat(arr, idx, M): # 初始化结果数组 res = np.zeros((*arr.shape[1:], M), dtype=arr.dtype) # 拓展索引维度适配广播规则 idx_expand = idx.reshape(-1, *([1]*(len(arr.shape)-1))) # 原地累加,自动处理重复索引 np.add.at(res, (Ellipsis, idx_expand), arr[..., np.newaxis]) return res
方案3:基于矩阵乘法的实现
该方案依赖BLAS优化的矩阵乘法算子,适合大批次场景,也可以无缝迁移到CuPy等GPU计算框架获得更高性能:
def bucket_sum_matmul(arr, idx, M): B = arr.shape[0] # 拉平非批次维度 arr_flat = arr.reshape(B, -1) # 生成索引的onehot编码矩阵 onehot = np.zeros((B, M), dtype=arr.dtype) onehot[np.arange(B), idx] = 1 # 矩阵乘法完成分桶求和 res_flat = arr_flat.T @ onehot # 恢复目标维度 return res_flat.reshape((*arr.shape[1:], M))
注意:以上三种方案的计算结果完全一致,相比朴素for循环性能提升可达10~1000倍,具体取决于批次大小、数组维度和分桶数量。
内容的提问来源于stack exchange,提问作者Jav
相关产品推荐
相关产品推荐

