PyTorch torch.nn.AvgPool2d Function
PyTorch torch.nn Reference Manual
torch.nn.AvgPool2dIt is a 2D average pooling module in PyTorch.
It computes the average for each window of the input, commonly used for downsampling and feature aggregation.
Function Definition
torch.nn.AvgPool2d(kernel_size, stride=None, padding=0, ceil_mode=False, count_include_pad=True)
Parameter Description
kernel_size: pooling window sizestride: stride, defaults to kernel_sizepadding: padding sizeceil_mode: whether to use ceil to compute output size
Usage Examples
Example 1: Basic Usage
Example
import torch
import torch.nn as nn
pool = nn.AvgPool2d(kernel_size=2, stride=2)
x = torch.randn(1, 1, 4, 4)
output = pool(x)
print("Input:n", x.squeeze().tolist())
print("nOutput:n", output.squeeze().tolist())
print("Shape:", x.shape, "->", output.shape)
import torch.nn as nn
pool = nn.AvgPool2d(kernel_size=2, stride=2)
x = torch.randn(1, 1, 4, 4)
output = pool(x)
print("Input:n", x.squeeze().tolist())
print("nOutput:n", output.squeeze().tolist())
print("Shape:", x.shape, "->", output.shape)
Example 2: Global Average Pooling
Example
import torch
import torch.nn as nn
# Global average pooling
gap = nn.AdaptiveAvgPool2d(1)
x = torch.randn(4, 64, 16, 16)
out = gap(x)
print("Input:", x.shape)
print("Output:", out.shape)
print("Number of averages per channel:", out.numel() // 64)
import torch.nn as nn
# Global average pooling
gap = nn.AdaptiveAvgPool2d(1)
x = torch.randn(4, 64, 16, 16)
out = gap(x)
print("Input:", x.shape)
print("Output:", out.shape)
print("Number of averages per channel:", out.numel() // 64)
Example 3: Comparison with MaxPool
Example
import torch
import torch.nn as nn
x = torch.tensor([[[
[1, 2, 3, 4],
[5, 6, 7, 8],
[9, 10, 11, 12],
[13, 14, 15, 16]
]]], dtype=torch.float32)
maxpool = nn.MaxPool2d(2, 2)
avgpool = nn.AvgPool2d(2, 2)
print("Input:n", x[0, 0])
print("nMaxPool:n", maxpool(x)[0, 0])
print("nAvgPool:n", avgpool(x)[0, 0])
import torch.nn as nn
x = torch.tensor([[[
[1, 2, 3, 4],
[5, 6, 7, 8],
[9, 10, 11, 12],
[13, 14, 15, 16]
]]], dtype=torch.float32)
maxpool = nn.MaxPool2d(2, 2)
avgpool = nn.AvgPool2d(2, 2)
print("Input:n", x[0, 0])
print("nMaxPool:n", maxpool(x)[0, 0])
print("nAvgPool:n", avgpool(x)[0, 0])
Use Cases
- Feature aggregation: Reduce spatial dimensions
- Global average pooling: Replace fully connected layers
- Smooth featuresReduce noise
Other extensions