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

使用Numpy实现Max Pooling时遇数组重塑尺寸错误求助

Numpy实现Max Pooling的维度错误修复

错误原因分析

你遇到的无法将尺寸为2883的数组重塑为(2,3,31,31)错误,核心问题如下:

  1. 批量维度未处理:代码仅计算了第一个样本(x[0])的池化结果,但输入是包含2个样本的张量,导致实际元素数(3×31×31=2883)远小于目标形状所需的元素数(2×3×31×31=5766)。
  2. 循环逻辑冗余:通过range(0, x_height, stride)加判断的方式遍历窗口,未利用已计算好的y_height和y_width直接控制循环次数,易出现计数偏差。
  3. 输出维度处理错误:np.hsplit(output,1)完全多余,反而引入不必要的维度,导致后续重塑失败。
  4. 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 18:35:26