跳转到内容
新建笔记

自适应平均池化:窗口推导与 PyTorch 对照实现

自适应平均池化直接指定输出的空间尺寸,再根据输入尺寸决定每个输出位置覆盖的区域。它保留通道数;“2D”指两个空间轴,不是说输入只有两个维度。

实际模型优先使用 nn.AdaptiveAvgPool2d 或 torch.nn.functional.adaptive_avg_pool2d。下面的实现用于理解边界和梯度,不作为内置算子的性能替代。

输入可以是 [N, C, H, W],也可以是不带批次的 [C, H, W]。输出空间尺寸可以写成整数 s,表示 (s, s),或写成 (H_out, W_out);其中一项为 None 时,保留该输入轴的大小。官方接口

例子输入形状输出形状
AdaptiveAvgPool2d(1)[N, C, H, W][N, C, 1, 1],每个通道取全局平均
AdaptiveAvgPool2d((3, 2))[N, C, 5, 7][N, C, 3, 2]
AdaptiveAvgPool2d((None, 2))[C, 5, 7][C, 5, 2]

输出尺寸也可以大于输入尺寸。此时多个输出位置可能依赖相同输入元素;这不是学习出了新的细节,也不能把它等同于双线性插值。

2. 用整数边界定义平均区域

跳转到“2. 用整数边界定义平均区域”

对输出行 i=0,…,Hout−1i=0,\ldots,H_{\mathrm{out}}-1,输入区间为左闭右开的

hs(i)=⌊iHinHout⌋,he(i)=⌈(i+1)HinHout⌉.h_s(i)=\left\lfloor\frac{iH_{\mathrm{in}}}{H_{\mathrm{out}}}\right\rfloor, \qquad h_e(i)=\left\lceil\frac{(i+1)H_{\mathrm{in}}}{H_{\mathrm{out}}}\right\rceil.

列方向同理得到 ws(j)w_s(j)、we(j)w_e(j)。每个输出元素对该矩形求平均:

yn,c,i,j=∑h=hs(i)he(i)−1∑w=ws(j)we(j)−1xn,c,h,w(he(i)−hs(i))(we(j)−ws(j)).y_{n,c,i,j} =\frac{\displaystyle\sum_{h=h_s(i)}^{h_e(i)-1} \sum_{w=w_s(j)}^{w_e(j)-1}x_{n,c,h,w}} {(h_e(i)-h_s(i))(w_e(j)-w_s(j))}.

例如长度 5 变成 3,三个区间是 [0:2]、[1:4]、[3:5]。它们既可能大小不同,也可能重叠,不能概括为“把输入平均切成互不重叠的几块”。边界定义可对照 PyTorch 2.8 的 AdaptivePooling 辅助函数。

对正整数,向上取整可写成整数运算 (a + b - 1) // b,因此不必为这些下标引入 NumPy 或浮点除法。

3. 修正原示例的几个约束

跳转到“3. 修正原示例的几个约束”

原实现要求调用者另外传入 input_size,却实际从 x 取数据。一旦二者不一致,索引区域就会错误。尺寸应该来自当前输入的 x.shape[-2:]。

原实现还在每个循环中调用 x.float(),会把输入转成 float32。这样可能改变数值精度,也不能保留 float64 梯度检查需要的类型。求平均和拼接本身支持自动求导,不需要通过类型转换“开启”梯度。

以下教学实现明确只接受常规稠密、非空的三维或四维 float32 / float64 张量,支持整数输出尺寸及二元尺寸中的 None。这比内置接口的完整设备、类型和空维支持范围更窄。

4. 可运行实现与前向、反向对照

跳转到“4. 可运行实现与前向、反向对照”
import torch
from torch import nn
from torch.nn import functional as F
class CustomAdaptiveAvgPool2D(nn.Module):
def __init__(self, output_size):
super().__init__()
if isinstance(output_size, int) and not isinstance(output_size, bool):
output_size = (output_size, output_size)
if not isinstance(output_size, (tuple, list)) or len(output_size) != 2:
raise ValueError("output_size must be an int or a pair")
for size in output_size:
if size is not None and (
not isinstance(size, int) or isinstance(size, bool) or size <= 0
):
raise ValueError("output sizes must be positive integers or None")
self.output_size = tuple(output_size)
def forward(self, x):
if not isinstance(x, torch.Tensor) or x.ndim not in (3, 4):
raise ValueError("expected a CHW or NCHW tensor")
if x.layout != torch.strided:
raise ValueError("this teaching example requires dense strided input")
if x.dtype not in (torch.float32, torch.float64):
raise TypeError("this teaching example supports float32 and float64")
if x.numel() == 0:
raise ValueError("this teaching example requires nonempty input")
height, width = x.shape[-2:]
out_height = height if self.output_size[0] is None else self.output_size[0]
out_width = width if self.output_size[1] is None else self.output_size[1]
rows = []
for i in range(out_height):
hs = i * height // out_height
he = ((i + 1) * height + out_height - 1) // out_height
columns = []
for j in range(out_width):
ws = j * width // out_width
we = ((j + 1) * width + out_width - 1) // out_width
cell = x[..., hs:he, ws:we].mean(dim=(-2, -1), keepdim=True)
columns.append(cell)
rows.append(torch.cat(columns, dim=-1))
return torch.cat(rows, dim=-2)
# 长度5到3:重叠区间的均值为0.5、2、3.5。
line = torch.arange(5, dtype=torch.float64).reshape(1, 1, 5, 1)
result = CustomAdaptiveAvgPool2D((3, 1))(line)
torch.testing.assert_close(result.flatten(), torch.tensor([0.5, 2.0, 3.5],
dtype=torch.float64))
# 测试矩形、升采样、无批次、None维度与全局平均。
torch.manual_seed(11)
cases = [((2, 3, 5, 7), (3, 2)), ((1, 2, 2, 3), (5, 4)),
((3, 5, 7), (None, 2)), ((2, 3, 4, 6), 1)]
for shape, output_size in cases:
data = torch.randn(shape, dtype=torch.float64)
custom_input = data.clone().requires_grad_()
native_input = data.clone().requires_grad_()
custom = CustomAdaptiveAvgPool2D(output_size)(custom_input)
native = F.adaptive_avg_pool2d(native_input, output_size)
assert custom.dtype == data.dtype and custom.device == data.device
torch.testing.assert_close(custom, native, rtol=1e-12, atol=1e-12)
upstream = torch.randn_like(custom)
custom_grad, = torch.autograd.grad(custom, custom_input, upstream)
native_grad, = torch.autograd.grad(native, native_input, upstream)
torch.testing.assert_close(custom_grad, native_grad, rtol=1e-12, atol=1e-12)
print("自适应平均池化的输出与梯度检查通过")

平均区域发生重叠时,一个输入元素会向多个输出位置贡献数值;反向传播要将这些区域的梯度贡献相加。只对照输出形状或前向值,还不足以验证自定义算子用于训练时的行为。

如果只是固定卷积网络的输出特征尺寸,直接使用内置层。如果是在研究新算子,先用小整数输入检查区域,再用随机 float64 输入对照前向与梯度,并验证不整除、输出大于输入、非连续输入以及非法尺寸。

这个逐格 Python 循环会创建较多中间结果,速度和内存开销不能由功能等价推断。本文不声称它能直接替代内置实现进行导出、编译优化或高吞吐推理。

张量视图与形状规则见张量基础;池化在网络中的位置见torch.nn 模块。