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

PyTorch中通道维度2D最大池化实现及尺寸匹配问题

问题分析与解决方案

原代码的思路方向是对的:利用3D池化同时覆盖通道维度和空间维度的池化需求,但因为空间维度(H/W)的池化核尺寸为2、步长为1,无法通过整数padding直接让输出尺寸与输入一致(计算可知需要padding=0.5,PyTorch不支持非整数padding),所以导致输出尺寸不符合预期。

正确实现方式

方式一:使用Unfold提取窗口后做全局最大池化

这种方式更直观,精准匹配需求:提取每个空间位置的2x2邻域窗口,再对窗口内所有通道取最大值。

import torch
import torch.nn as nn

torch.manual_seed(0)

B, CH, H, W = 8, 64, 128, 128
x_batch = torch.randn((B, CH, H, W))

# 提取2x2窗口,padding=1保证边缘位置也能取到2x2窗口,步长1对应每个空间位置
unfold = nn.Unfold(kernel_size=(2,2), stride=1, padding=1)
x_unfold = unfold(x_batch)  # 形状: (B, CH*2*2, H*W)

# 对所有窗口内的通道值取最大
x_max = x_unfold.max(dim=1, keepdim=True)[0]
# 重塑回目标形状
x_max = x_max.view(B, 1, H, W)

print(x_max.shape)  # 输出: torch.Size([8, 1, 128, 128])

方式二:3D池化后裁剪多余维度

如果坚持用3D池化,可以先通过padding得到稍大的输出,再裁剪掉多余的行和列:

import torch
import torch.nn as nn

torch.manual_seed(0)

B, CH, H, W = 8, 64, 128, 128
x_batch = torch.randn((B, CH, H, W))

# 用padding=(0,1,1)得到(8,1,129,129)的输出
max3d = nn.MaxPool3d((64,2,2), stride=1, padding=(0,1,1))
x_max = max3d(x_batch)
# 裁剪掉最后一行和最后一列,得到目标尺寸
x_max = x_max[:, :, :-1, :-1]

print(x_max.shape)  # 输出: torch.Size([8, 1, 128, 128])

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 03:46:04