如何在不使用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的值:
步骤说明
- 翻转kernel:保持和原代码一致的预处理逻辑(如果你的需求是卷积操作,翻转是必要的;若只是单纯叠加kernel,可跳过此步骤)
- 生成偏移量:计算kernel内每个元素相对于中心的坐标偏移
- 计算所有更新坐标:基于每个非零元素的位置,结合偏移量得到所有需要叠加kernel值的输出数组坐标
- 过滤边界无效坐标:排除那些超出输入数组范围的坐标,避免索引错误
- 批量累加:使用
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
相关产品推荐
相关产品推荐

