使用Numpy实现Max Pooling时遇数组重塑尺寸错误求助
Numpy实现Max Pooling的维度错误修复
错误原因分析
你遇到的无法将尺寸为2883的数组重塑为(2,3,31,31)错误,核心问题如下:
- 批量维度未处理:代码仅计算了第一个样本(
x[0])的池化结果,但输入是包含2个样本的张量,导致实际元素数(3×31×31=2883)远小于目标形状所需的元素数(2×3×31×31=5766)。 - 循环逻辑冗余:通过
range(0, x_height, stride)加判断的方式遍历窗口,未利用已计算好的y_height和y_width直接控制循环次数,易出现计数偏差。 - 输出维度处理错误:
np.hsplit(output,1)完全多余,反而引入不必要的维度,导致后续重塑失败。 - padding参数未实现:函数定义了padding参数,但未在代码中添加零填充逻辑。
修正后的代码
import numpy as np def maxpool(x, kernel_size, stride, padding): """ Args: x: numpy array with size (N, C_in, H_in, W_in), kernel_size: size of the window to take a max over, stride: stride of the window, padding: implicit zero padding to be added on both sides, Return: y: numpy array of size (N, C_out, H_out, W_out). """ # 转换为numpy数组(兼容torch tensor输入) x = np.array(x) N, C_in, H_in, W_in = x.shape # 计算输出维度,考虑padding H_out = (H_in + 2 * padding - kernel_size) // stride + 1 W_out = (W_in + 2 * padding - kernel_size) // stride + 1 # 对输入进行零填充 if padding > 0: x = np.pad(x, ((0,0), (0,0), (padding,padding), (padding,padding)), mode='constant') # 初始化输出数组 output = np.zeros((N, C_in, H_out, W_out)) # 遍历每个样本、通道、输出位置 for n in range(N): for c in range(C_in): for h in range(H_out): for w in range(W_out): # 计算当前池化窗口的起始位置 h_start = h * stride w_start = w * stride # 提取窗口区域并取最大值 window = x[n, c, h_start:h_start+kernel_size, w_start:w_start+kernel_size] output[n, c, h, w] = np.max(window) return output
关键改进点
- 批量处理:新增对
N个样本的循环,确保所有输入样本都被计算。 - padding实现:使用
np.pad对输入的高度和宽度两侧添加零填充,匹配参数定义的功能。 - 直接窗口定位:通过
H_out和W_out控制循环次数,直接计算每个输出位置对应的窗口起始坐标,避免冗余判断。 - 高效取最大值:利用
np.max直接计算窗口最大值,替代手动拆分和遍历的低效方式。
测试验证
对于输入x = np.random.randn(2,3,32,32),调用maxpool(x, kernel_size=2, stride=1, padding=0),输出形状为(2,3,31,31),符合预期。
内容的提问来源于stack exchange,提问作者Steamer
相关产品推荐
相关产品推荐

