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

如何高效实现二维矩阵指定列的原地累加(列索引为二维数组)

高效实现带重复列索引的二维矩阵原地累加操作

问题背景

需要对二维矩阵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专门为重复索引原地累加设计的函数,会遍历所有索引位置执行累加,而非覆盖重复索引的赋值。实现步骤如下:

  1. 将二维的col_idx展平为一维数组,得到所有需要累加的列索引序列;
  2. 将values调整为与展平后索引匹配的形状,确保每个索引对应正确的累加值;
  3. 调用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 18:15:54