如何仅在指定位置执行二维卷积以降低计算与内存开销
单点卷积值实现方案
不需要调用全图卷积类的内置函数,直接基于卷积的计算逻辑取对应区域计算即可,计算量仅和卷积核大小相关,和输入数据的整体尺寸无关,完全适配N、M为10000级别的场景,内存和计算开销都可以压到最低。
卷积的本质就是卷积核(或互相关核,取决于你遵循的卷积定义)在对应感受野内和输入逐元素相乘后求和,要拿单个坐标的输出值,根本不需要遍历全图所有位置,只需要取出该坐标对应的输入感受野切片,和卷积核做乘加计算就行。
具体实现逻辑
首先确认两个基础配置:
- 卷积的锚点位置:常规3x3卷积的锚点在核的中心,也就是
kernel[1,1]位置对齐你要计算的目标坐标点 - padding规则:分为
valid(不补边,越界点无输出)和same(补0/补指定值,保证输出和输入尺寸一致)两种,按需选择即可
代码实现
import numpy as np def calc_single_conv_point(data, kernel, target_x, target_y, pad_mode="same", flip_kernel=False): """ 计算输入数据上单个坐标点的卷积输出值 :param data: 二维输入数组,shape (H,W) :param kernel: 二维卷积核,shape (kH,kW) :param target_x: 目标点x坐标(行号) :param target_y: 目标点y坐标(列号) :param pad_mode: padding模式,支持valid/same,same模式默认补0 :param flip_kernel: 是否翻转卷积核,传统数字图像处理卷积需设为True,深度学习框架的Conv2d(互相关)设为False即可 """ k_h, k_w = kernel.shape anchor_h, anchor_w = k_h // 2, k_w // 2 # 计算感受野在原数据上的坐标范围 h_start = target_x - anchor_h h_end = target_x + (k_h - anchor_h) w_start = target_y - anchor_w w_end = target_y + (k_w - anchor_w) if pad_mode == "valid": if h_start < 0 or h_end > data.shape[0] or w_start <0 or w_end > data.shape[1]: raise ValueError("目标坐标超出valid卷积的合法输出范围") data_patch = data[h_start:h_end, w_start:w_end] elif pad_mode == "same": data_patch = np.zeros((k_h, k_w), dtype=data.dtype) # 计算原数据和patch的重叠区域坐标 p_h_st = max(0, -h_start) p_h_ed = k_h - max(0, h_end - data.shape[0]) p_w_st = max(0, -w_start) p_w_ed = k_w - max(0, w_end - data.shape[1]) d_h_st = max(0, h_start) d_h_ed = min(data.shape[0], h_end) d_w_st = max(0, w_start) d_w_ed = min(data.shape[1], w_end) data_patch[p_h_st:p_h_ed, p_w_st:p_w_ed] = data[d_h_st:d_h_ed, d_w_st:d_w_ed] else: raise ValueError("仅支持valid、same两种padding模式") # 按需翻转卷积核 if flip_kernel: kernel = np.flip(kernel, axis=(0,1)) return np.sum(data_patch * kernel) # 调用示例:计算(10,37)位置的卷积值 N = 10000 data = np.random.rand(N, N) kernel = np.random.rand(3, 3) point_val = calc_single_conv_point(data, kernel, 10, 37, pad_mode="same", flip_kernel=False)
性能说明
- 该实现的计算复杂度为O(M²),和输入数据的尺寸N完全无关,哪怕N达到10^5级别,计算耗时也只和卷积核大小有关
- 内存开销仅为M*M尺寸的临时数组,不会加载全量中间结果,完全规避全图卷积的冗余计算
- 所有切片、乘加操作都是numpy底层C实现,性能和内置函数没有差异,不需要额外引入其他依赖
内容的提问来源于stack exchange,提问作者deltasata
相关产品推荐
相关产品推荐

