如何高效实现二维矩阵指定列的原地累加(列索引为二维数组)
高效实现带重复列索引的二维矩阵原地累加操作
问题背景
需要对二维矩阵grad的指定列执行**原地累加(+=)**操作,列索引由二维数组col_idx提供(其中存在重复索引)。当前嵌套循环实现的loop_fun结果正确但效率极低;直接用向量级别的+=操作会因重复索引仅保留最后一次赋值,导致结果错误(如loop_fun_wrong所示)。希望借助np.add.at实现高效的正确版本。
代码示例
import numpy as np np.random.seed(1) N, V, D, B = 5, 25000, 300, 512 grad = np.zeros((D, V)) # col_idx包含重复索引,直接向量级+=会丢失重复索引的累加操作 col_idx = np.random.randint(0, V-1, size=(N, B)) # 强制构造一组重复索引示例 col_idx[:, 0] = np.array([0, 100, V-1, 0, V-1]) values = np.random.normal(size=(D, B)) def loop_fun(grad_mat, col_mat, val_mat): for b in range(col_mat.shape[1]): for row in range(col_mat.shape[0]): grad_mat[:, col_mat[row, b]] += val_mat[:, b] return grad_mat def loop_fun_wrong(grad_mat, col_mat, val_mat): # 错误:重复索引仅保留最后一次赋值 for b in range(col_mat.shape[1]): grad_mat[:, col_mat[:, b]] += val_mat[:, b, np.newaxis] return grad_mat grad = loop_fun(grad, col_idx, values) grad_wrong = loop_fun_wrong(np.zeros_like(grad), col_idx, values) print(f'{np.allclose(grad, grad_wrong)=}') # False
解决方案:使用np.add.at实现向量化累加
np.add.at是numpy专门为重复索引原地累加设计的函数,会遍历所有索引位置执行累加,而非覆盖重复索引的赋值。实现步骤如下:
- 将二维的
col_idx展平为一维数组,得到所有需要累加的列索引序列; - 将
values调整为与展平后索引匹配的形状,确保每个索引对应正确的累加值; - 调用
np.add.at完成原地累加。
实现代码
基础版本
def vectorized_fun(grad_mat, col_mat, val_mat): # 展平列索引为一维 flat_col_idx = col_mat.flatten() # 将values沿列方向重复N次(匹配col_mat的行数),得到(D, N*B)形状的数组 repeated_vals = np.repeat(val_mat, repeats=col_mat.shape[0], axis=1) # 执行原地累加:slice(None)表示选中所有行,flat_col_idx指定列索引 np.add.at(grad_mat, (slice(None), flat_col_idx), repeated_vals) return grad_mat
内存优化版本
避免np.repeat创建大数组,通过广播调整形状后直接展平,节省内存:
def vectorized_fun_memory_efficient(grad_mat, col_mat, val_mat): D = grad_mat.shape[0] # 将val_mat调整为(D,1,B),与col_mat的(1,N,B)广播后展平为(D, N*B) flat_vals = val_mat[:, np.newaxis, :].reshape(D, -1) # 展平列索引并执行累加 np.add.at(grad_mat, (slice(None), col_mat.ravel()), flat_vals) return grad_mat
正确性验证
grad_vectorized = vectorized_fun(np.zeros_like(grad), col_idx, values) print(f'{np.allclose(grad, grad_vectorized)=}') # True grad_vectorized_eff = vectorized_fun_memory_efficient(np.zeros_like(grad), col_idx, values) print(f'{np.allclose(grad, grad_vectorized_eff)=}') # True
效率说明
np.add.at基于numpy底层的C实现,完全避免了Python层面的嵌套循环,在大数组场景下(如示例中V=25000、B=512),速度会比嵌套循环提升几个数量级。
内容的提问来源于stack exchange,提问作者dule arnaux
相关产品推荐
相关产品推荐

