如何仅用Numpy实现PyTorch的grid_sample函数?
仅用Numpy实现PyTorch风格的grid_sample函数
核心思路
PyTorch的grid_sample默认采用双线性插值,核心是将归一化到[-1,1]的网格坐标转换为图像像素坐标,然后通过邻域像素的加权求和得到输出像素值。以下是完全基于Numpy的实现方案,严格对齐PyTorch的默认行为(mode='bilinear'、padding_mode='constant'、align_corners=True)。
实现代码
import numpy as np def numpy_grid_sample(img, grid, mode='bilinear', padding_mode='constant', align_corners=True): """ 仅用Numpy实现PyTorch风格的grid_sample 参数: img: 输入图像,形状为(N, C, H, W),N=批量数,C=通道数,H=高度,W=宽度 grid: 变换网格,形状为(N, H_out, W_out, 2),最后一维为(x, y)归一化坐标(范围[-1,1]) mode: 插值模式,仅支持'billinear'(默认)和'nearest' padding_mode: 边界填充模式,仅支持'constant'(默认,填充0)和'border'(填充边界像素) align_corners: 是否对齐角落像素,对应PyTorch同名参数 返回: output: 变换后的图像,形状为(N, C, H_out, W_out) """ N, C, H, W = img.shape N_out, H_out, W_out, _ = grid.shape # 将归一化坐标转换为像素坐标 if align_corners: x = (grid[..., 0] + 1) * (W - 1) / 2 y = (grid[..., 1] + 1) * (H - 1) / 2 else: x = (grid[..., 0] + 1) * W / 2 - 0.5 y = (grid[..., 1] + 1) * H / 2 - 0.5 if mode == 'nearest': # 最近邻插值:直接取最近的整数坐标 x_idx = np.around(x).astype(np.int32) y_idx = np.around(y).astype(np.int32) # 处理边界 if padding_mode == 'constant' or padding_mode == 'border': x_idx = np.clip(x_idx, 0, W-1) y_idx = np.clip(y_idx, 0, H-1) # 构造批量索引,实现批量维度的广播 batch_idx = np.arange(N)[:, np.newaxis, np.newaxis] output = img[batch_idx, :, y_idx, x_idx] return output elif mode == 'bilinear': # 双线性插值:计算四个邻域像素坐标 x0 = np.floor(x).astype(np.int32) x1 = x0 + 1 y0 = np.floor(y).astype(np.int32) y1 = y0 + 1 # 处理边界填充 if padding_mode == 'constant' or padding_mode == 'border': x0 = np.clip(x0, 0, W-1) x1 = np.clip(x1, 0, W-1) y0 = np.clip(y0, 0, H-1) y1 = np.clip(y1, 0, H-1) # 计算四个邻域的插值权重 wa = (x1 - x) * (y1 - y) wb = (x1 - x) * (y - y0) wc = (x - x0) * (y1 - y) wd = (x - x0) * (y - y0) # 扩展权重维度,适配通道维度的广播 wa = np.expand_dims(wa, axis=1) wb = np.expand_dims(wb, axis=1) wc = np.expand_dims(wc, axis=1) wd = np.expand_dims(wd, axis=1) # 构造批量索引,提取四个邻域的像素值 batch_idx = np.arange(N)[:, np.newaxis, np.newaxis] img_a = img[batch_idx, :, y0, x0] img_b = img[batch_idx, :, y1, x0] img_c = img[batch_idx, :, y0, x1] img_d = img[batch_idx, :, y1, x1] # 加权求和得到最终输出 output = wa * img_a + wb * img_b + wc * img_c + wd * img_d return output else: raise ValueError(f"不支持的插值模式: {mode}")
关键细节说明
- 坐标转换:严格对应PyTorch的
align_corners参数逻辑,确保坐标映射完全一致。 - 批量处理:通过
batch_idx实现批量维度的索引广播,避免循环处理每个样本。 - 边界处理:用
np.clip限制坐标范围,实现constant和border填充模式,超出图像范围的像素会被截断到边界(constant模式下超出部分最终权重为0,等价于填充0)。 - 插值逻辑:双线性插值的权重计算完全遵循数学公式,保证结果和PyTorch对齐。
验证方法
可以将Numpy实现的结果与PyTorch原生grid_sample对比,误差应接近0:
import torch # 构造测试数据 np_img = np.random.rand(1, 3, 20, 20).astype(np.float32) torch_img = torch.from_numpy(np_img) # 构造测试变换网格(替换为你已实现的Numpy版affine_grid) theta = np.array([[1, 0, 0.2], [0, 1, 0.1]], dtype=np.float32)[np.newaxis, ...] np_grid = your_numpy_affine_grid(theta, (1, 3, 20, 20)) torch_grid = torch.nn.functional.affine_grid(torch.from_numpy(theta), torch_img.shape, align_corners=True) # 计算结果 np_output = numpy_grid_sample(np_img, np_grid, align_corners=True) torch_output = torch.nn.functional.grid_sample(torch_img, torch_grid, align_corners=True).numpy() # 输出最大误差 print(f"与PyTorch结果的最大误差: {np.max(np.abs(np_output - torch_output)):.6f}")
内容的提问来源于stack exchange,提问作者pepe calero
相关产品推荐
相关产品推荐

