自适应平均池化直接指定输出的空间尺寸,再根据输入尺寸决定每个输出位置覆盖的区域。它保留通道数;“2D”指两个空间轴,不是说输入只有两个维度。
实际模型优先使用 nn.AdaptiveAvgPool2d 或 torch.nn.functional.adaptive_avg_pool2d。下面的实现用于理解边界和梯度,不作为内置算子的性能替代。
1. 输入与输出的约定
跳转到“1. 输入与输出的约定”输入可以是 [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. 用整数边界定义平均区域”对输出行 ,输入区间为左闭右开的
列方向同理得到 、。每个输出元素对该矩形求平均:
例如长度 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 torchfrom torch import nnfrom 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("自适应平均池化的输出与梯度检查通过")平均区域发生重叠时,一个输入元素会向多个输出位置贡献数值;反向传播要将这些区域的梯度贡献相加。只对照输出形状或前向值,还不足以验证自定义算子用于训练时的行为。
5. 选择与验证
跳转到“5. 选择与验证”如果只是固定卷积网络的输出特征尺寸,直接使用内置层。如果是在研究新算子,先用小整数输入检查区域,再用随机 float64 输入对照前向与梯度,并验证不整除、输出大于输入、非连续输入以及非法尺寸。
这个逐格 Python 循环会创建较多中间结果,速度和内存开销不能由功能等价推断。本文不声称它能直接替代内置实现进行导出、编译优化或高吞吐推理。
张量视图与形状规则见张量基础;池化在网络中的位置见torch.nn 模块。