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

如何在不使用for循环的情况下高效修改稀疏Numpy数组子块?

问题

我有一个仅包含0或1的稀疏numpy数组,还有一个k×k维度的kernel。需要针对数组中的每个非零元素,以其为中心截取k×k的子块,并将kernel叠加到该子块上。希望实现高效处理,尽可能避免使用for循环。

目前我尝试的代码如下:

import numpy as np

def sparse_convolutionv2(input, kernel):
    kernel = np.flipud(np.fliplr(kernel))
    
    # 获取非零元素索引
    non_zero_indices = np.transpose(np.nonzero(input))
    
    # 计算子块的起止坐标
    row_start, row_end = non_zero_indices[:, 0] - kernel.shape[0] // 2, non_zero_indices[:, 0] + kernel.shape[0] // 2 + 1
    col_start, col_end = non_zero_indices[:, 1] - kernel.shape[1] // 2, non_zero_indices[:, 1] + kernel.shape[1] // 2 + 1
    output = np.zeros_like(input)

    # 叠加kernel
    for i in range(len(row_start)):
        output[row_start[i]:row_end[i], col_start[i]:col_end[i]] = np.add(output[row_start[i]:row_end[i], col_start[i]:col_end[i]], kernel)
    return output
高效无循环实现方案

利用numpy的向量化操作和np.add.at函数可以实现无循环的高效处理,核心思路是一次性生成所有需要更新的坐标,再批量累加kernel的值:

步骤说明

  1. 翻转kernel:保持和原代码一致的预处理逻辑(如果你的需求是卷积操作,翻转是必要的;若只是单纯叠加kernel,可跳过此步骤)
  2. 生成偏移量:计算kernel内每个元素相对于中心的坐标偏移
  3. 计算所有更新坐标:基于每个非零元素的位置,结合偏移量得到所有需要叠加kernel值的输出数组坐标
  4. 过滤边界无效坐标:排除那些超出输入数组范围的坐标,避免索引错误
  5. 批量累加:使用np.add.at处理重复坐标的累加,完成kernel的叠加

实现代码

import numpy as np

def sparse_convolution_vectorized(input_arr, kernel):
    # 预处理:翻转kernel(若无需卷积翻转可删除此行)
    kernel = np.flipud(np.fliplr(kernel))
    kh, kw = kernel.shape
    h_half = kh // 2
    w_half = kw // 2
    
    # 获取非零元素的行、列索引
    r, c = np.nonzero(input_arr)
    num_nonzero = len(r)
    
    # 生成kernel内所有元素的坐标偏移
    dr = np.arange(-h_half, h_half + 1)
    dc = np.arange(-w_half, w_half + 1)
    offsets = np.array(np.meshgrid(dr, dc)).reshape(2, -1).T  # 形状为(kh*kw, 2)
    
    # 计算所有需要更新的输出坐标
    output_r = r[:, None] + offsets[:, 0]  # 形状(num_nonzero, kh*kw)
    output_c = c[:, None] + offsets[:, 1]
    
    # 拉平坐标和kernel值
    output_r_flat = output_r.flatten()
    output_c_flat = output_c.flatten()
    kernel_flat = np.tile(kernel.flatten(), num_nonzero)
    
    # 过滤掉超出输入数组范围的无效坐标
    valid_mask = (output_r_flat >= 0) & (output_r_flat < input_arr.shape[0]) & \
                 (output_c_flat >= 0) & (output_c_flat < input_arr.shape[1])
    valid_r = output_r_flat[valid_mask]
    valid_c = output_c_flat[valid_mask]
    valid_kernel_vals = kernel_flat[valid_mask]
    
    # 批量累加更新输出数组
    output = np.zeros_like(input_arr)
    np.add.at(output, (valid_r, valid_c), valid_kernel_vals)
    
    return output

优势说明

  • 无循环:完全依赖numpy的向量化操作,避免了Python级别的for循环,大幅提升处理速度(尤其当非零元素数量较多时)
  • 高效累加:np.add.at专门用于处理重复索引的累加操作,完美适配多个kernel叠加到同一区域的场景
  • 边界处理:通过掩码过滤无效坐标,避免了索引越界问题

内容的提问来源于stack exchange,提问作者xxsilmarilionxx

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 04:45:35