PyTorch中如何纯张量实现带双线性插值的小数偏移图像平移
结论
这个功能完全可以仅通过PyTorch原生张量操作实现,不需要额外开发C扩展。PyTorch自带的张量算子已经覆盖了双线性插值、坐标映射、边界填充的全部需求,运行效率和手写C扩展没有量级差距,同时支持自动微分,可以直接嵌入模型训练流程。
实现逻辑
首先明确坐标映射规则:平移后的输出张量尺寸和输入保持[H,W]不变,输出位置(out_y, out_x)对应的原始输入坐标为:
- 原始水平坐标:
out_x - mu_x - 原始垂直坐标:
out_y - mu_y
以mu_x=5、mu_y=3的整数平移场景为例:输出位置(3,5)刚好对应输入的(0,0)位置,输出位置(0~2, 任意)、(任意, 0~4)对应的原始坐标都在输入边界外,直接填充0,完全匹配需求。
实现路径有两种,均不需要C++扩展:
- 最优方案是直接调用PyTorch原生的
F.grid_sample接口:该接口是官方底层优化过的采样算子,原生支持双线性插值、边界零填充配置,CPU/GPU运行效率极高。使用时仅需要把坐标归一化到接口要求的[-1,1]区间,插值模式选bilinear、边界填充模式选zeros、设置align_corners=False对齐常规图像处理坐标规则即可。 - 也可以手动实现双线性插值逻辑:通过原生张量运算生成每个输出点对应的4个邻域整数坐标、计算双线性权重、用掩码过滤边界外的采样点、加权求和得到最终值,只是代码量更大,效率和
grid_sample没有本质差异。
可直接运行的代码实现
import torch import torch.nn.functional as F def translate_2d(x: torch.Tensor, mu_x: float, mu_y: float) -> torch.Tensor: """ 对输入张量做双线性插值平移,边界外补0 Args: x: 输入张量,支持两种格式: 1. 二维张量,形状[H, W],单通道二维数据 2. 四维张量,形状[N, C, H, W],批量多通道图像格式 mu_x: x方向(水平向右)平移量,支持小数 mu_y: y方向(垂直向下)平移量,支持小数 Returns: 平移后的张量,形状和输入完全一致 """ # 统一输入维度为[N,C,H,W]格式适配grid_sample接口 input_dim = x.dim() if input_dim == 2: x = x.unsqueeze(0).unsqueeze(0) N, C, H, W = x.shape # 生成输出位置的基础坐标网格 out_y = torch.arange(H, device=x.device, dtype=x.dtype) out_x = torch.arange(W, device=x.device, dtype=x.dtype) grid_y, grid_x = torch.meshgrid(out_y, out_x, indexing='ij') # 映射到输入张量的采样坐标 src_x = grid_x - mu_x src_y = grid_y - mu_y # 坐标归一化到grid_sample要求的[-1, 1]范围 src_x_norm = 2 * src_x / (W - 1) - 1 src_y_norm = 2 * src_y / (H - 1) - 1 # 拼接为接口要求的[N, H, W, 2]格式采样网格 grid = torch.stack([src_x_norm, src_y_norm], dim=-1).expand(N, -1, -1, -1) # 调用原生双线性采样,边界外自动补0 out = F.grid_sample( x, grid, mode='bilinear', padding_mode='zeros', align_corners=False ) # 还原输入原本的维度格式 if input_dim == 2: out = out.squeeze(0).squeeze(0) return out
效果说明
- 当mu_x、mu_y为整数时,函数输出和整数平移、裁切补0的结果完全一致
- 当mu_x、mu_y为小数时,自动通过双线性插值计算像素值,全程无Python侧循环,所有计算均走PyTorch底层优化的算子实现
内容的提问来源于stack exchange,提问作者lbwnb123
相关产品推荐
相关产品推荐

