torch.nn.functional.grid_sample在2D灰度图中的工作原理及疑问解答
关于PyTorch
F.grid_sample 的工作机制解惑 问题描述
编写了一段对灰度图做变换的PyTorch代码,但对F.grid_sample的工作机制有两处困惑:
import torch import numpy as np # Gray Scale Image image = torch.tensor([[[1, 2, 3, 4], [5, 6, 7, 8], [9, 10, 11, 12], [13, 14, 15, 16]]] ).unsqueeze(0).float() # Define a simple grid with some shifts and rotations grid_x, grid_y = torch.meshgrid(torch.arange(4), torch.arange(4)) grid_x = grid_x.float() grid_y = grid_y.float() new_locs = torch.stack([grid_x + 0.2 * torch.sin(grid_y), grid_y - 0.1 * torch.cos(grid_x)], dim=2).unsqueeze(0).float() # Warp the image using grid_sample import torch.nn.functional as F warped_image = F.grid_sample(image, new_locs, align_corners=True, mode='nearest')
困惑点:
- 调试发现
new_locs的值处于[-1,1]区间,而常规图像处理中像素坐标以(0,0)为左上角且均为非负值,无法理解输出像素位置为何会出现负值。 - 不清楚插值模式(比如示例中的最近邻插值)的作用,不明白
F.grid_sample为何需要使用插值。
解答
1. 为什么坐标会在[-1,1]区间?
F.grid_sample要求输入的网格坐标是归一化坐标,而非常规的像素索引。这么设计是为了适配任意尺寸的输入图像——不管图像是256x256还是512x512,坐标范围都固定在[-1,1],不用每次都根据图像宽高换算像素位置。
当align_corners=True时,坐标映射规则很明确:
- 原始图像左上角的像素坐标(0,0)对应归一化后的(-1,-1)
- 原始图像右下角的像素坐标(H-1, W-1)对应归一化后的(1,1)
- 所有中间像素的坐标按线性比例缩放到[-1,1]区间
你调试时看到new_locs落在这个区间,是因为代码里的偏移计算刚好让值处于该范围,但本质上grid_sample的输入必须是这个区间的坐标,超出部分会按padding_mode参数的设置做边界处理(比如填充、镜像等)。
2. 为什么需要插值模式?
grid_sample的核心是按网格坐标从原始图像采样像素,生成变换后的新图像。但网格坐标通常是连续的浮点值(比如你代码里grid_x + 0.2 * torch.sin(grid_y)会得到带小数的坐标),这些位置并不刚好落在原始图像的整数像素上,这时候就需要用插值来估算这个非整数位置的像素值。
举个实际例子:如果网格里的某个坐标指向原始图像的(2.3, 1.7),这个位置介于四个像素之间,插值算法就派上用场了:
- 最近邻插值:直接取离这个点最近的整数坐标(比如(2,2))的像素值,计算最快但可能有锯齿
- 双线性插值:用周围4个像素的加权平均计算,结果更平滑
- 双三次插值:用周围更多像素加权,精度更高但计算量更大
如果没有插值,面对浮点坐标时就无法准确获取对应的像素信息,插值就是为了处理这种连续坐标的采样需求,让变换后的图像符合预期效果。
内容的提问来源于stack exchange,提问作者Swakshar Deb
相关产品推荐
相关产品推荐

