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

如何用C++自主实现PyTorch adaptive_avg_pool2d及其底层原理

问题描述

尝试复现PyTorch中adaptive_avg_pool2d的运算行为时,发现手动用固定参数AvgPool2d实现的结果与官方输出存在差异,测试代码如下:

def test_pool():
    a = np.fromfile("in.bin", dtype=np.float32)
    a = np.reshape(a, [1, 12, 25, 25])
    a = torch.as_tensor(a)

    b = F.adaptive_avg_pool2d(a, [7, 7])
    print(b)
    print(b.shape)

    avg_pool = torch.nn.AvgPool2d([7, 7], [3, 3])
    c = avg_pool(a)
    print(c)
    print(c.shape)

待解决的两个核心问题:

  • PyTorch中adaptive_avg_pool2d的底层实现原理是什么
  • 如何基于原理在C++中自主实现该算子,得到与PyTorch官方完全一致的运算结果
底层实现原理

你用固定核大小、固定步长的AvgPool2d复现失败是必然结果:自适应平均池化不存在全局统一的卷积核尺寸与步长参数,它会逐输出位置动态计算对应的输入特征图切片范围,仅在输入输出尺寸满足特定整除关系时,才会等价于固定参数的平均池化。

PyTorch官方实现的计算逻辑完全对齐以下规则:
假设输入特征图单通道尺寸为(H, W),指定输出尺寸为(OH, OW),对于输出特征图上坐标为(i, j)的点(i取值范围0~OH-1,j取值范围0~OW-1):

  1. 计算高度方向对应的输入切片范围
    h_start = floor(i * H / OH)
    h_end = ceil((i + 1) * H / OH)
    h_k = h_end - h_start
    
  2. 计算宽度方向对应的输入切片范围
    w_start = floor(j * W / OW)
    w_end = ceil((j + 1) * W / OW)
    w_k = w_end - w_start
    
  3. 该输出点的值 = 输入特征图h_start到h_end行、w_start到w_end列围成的矩形区域内所有元素的平均值,即区域元素总和除以h_k * w_k。

以你测试用的输入尺寸H=W=25、输出尺寸OH=OW=7为例,不同输出位置对应的池化窗口大小是动态变化的:

  • 第0行输出:h_start=0,h_end=4,窗口高度为4
  • 第1行输出:h_start=3,h_end=8,窗口高度为5
    你代码里固定用7x7窗口、步长3的平均池化,每个窗口尺寸完全一致,和动态变化的窗口逻辑不匹配,结果自然存在差异。
C++实现对齐要点

要做到和PyTorch官方输出完全一致,实现时必须遵守以下规则:

  • 索引用整数运算实现,避免浮点计算带来的精度误差:上述floor和ceil逻辑可以直接用整数算术等价实现,和PyTorch内部写法完全一致:
    // 高度方向索引计算
    int h_start = (i * H) / OH; // C++整数除法天然向下取整,等价floor
    int h_end = ((i + 1) * H + OH - 1) / OH; // 整数向上取整写法,等价ceil
    // 宽度方向索引计算
    int w_start = (j * W) / OW;
    int w_end = ((j + 1) * W + OW - 1) / OW;
    
  • 累加过程保持和输入一致的精度:比如输入为float32类型时,累加和也用float32存储即可,不要随意转成float16或者double类型,避免累加误差。
  • 不需要额外做padding或越界裁剪:按照上述公式计算出的切片范围天然落在输入的合法索引区间内,不存在越界问题。
  • 多batch、多通道场景下,逐batch、逐通道独立计算上述池化逻辑即可,batch维和通道维的维度顺序和输入保持一致,不参与池化计算。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 16:18:21