神经网络框架中可微分forward mapping/warping实现方案咨询
实现可微分的正向图像扭曲(Forward Warping)
你提到的grid_sample是反向映射(backward warping)——从目标像素出发找对应的源像素,这也是为啥每个目标像素只关联源的局部邻域。而你需要的**正向映射(forward warping)**正好反过来:每个源像素根据偏移映射到目标图像的某个(可能非整数)位置,并且把该像素的数值按权重分配给目标的邻域像素,这样一个源像素就能影响目标图像中多个(甚至非相邻)的像素,同时整个过程保持可微分,完全适配神经网络训练。
下面是基于PyTorch的可微分正向映射实现,我会结合你的示例代码来修改:
核心思路
正向映射的关键是:
- 对源图像的每个像素,计算它在目标图像中的映射坐标
- 针对浮点型的映射坐标,用双线性插值计算该像素对目标邻域四个像素的贡献权重
- 用PyTorch的
scatter_add_操作累加所有源像素的贡献,这个操作是完全可微分的,能和神经网络的反向传播兼容
完整实现代码
import torch import torch.nn.functional as F import numpy as np import matplotlib.pyplot as plt # 创建人工图像 im = np.zeros((9,9)) im[4,4] = 1 # 图像扭曲参数:这里我们给源像素(4,4)一个偏移,让它映射到目标的(4,3)附近 u = np.zeros((9,9)) v = np.zeros((9,9)) u[4,4] += 1 # 注意:正向映射中,u是x方向的偏移,对应源像素的x坐标加上u得到目标x坐标 # 转换为PyTorch张量 h, w = im.shape im_t = torch.from_numpy(im).float().view(1, 1, h, w) # shape: (B, C, H, W) u_t = torch.from_numpy(u).float().view(1, h, w) v_t = torch.from_numpy(v).float().view(1, h, w) batch_size, channels, src_h, src_w = im_t.size() # -------------------------- 正向映射核心代码 -------------------------- # 1. 生成源图像每个像素的坐标 (x, y) x_src = torch.arange(src_w, dtype=torch.float32).repeat(src_h, 1).unsqueeze(0) # shape: (1, H, W) y_src = torch.arange(src_h, dtype=torch.float32).repeat(src_w, 1).t().unsqueeze(0) # shape: (1, H, W) # 2. 计算目标图像中的坐标 (x_dst, y_dst) = (x_src + u, y_src + v) x_dst = x_src + u_t y_dst = y_src + v_t # 3. 处理边界:把超出目标图像范围的坐标clamp到图像内(也可以直接丢弃,这里选clamp) x_dst = torch.clamp(x_dst, 0, src_w - 1 - 1e-6) # 减去1e-6避免取整到src_w y_dst = torch.clamp(y_dst, 0, src_h - 1 - 1e-6) # 4. 计算双线性插值的四个邻域坐标和权重 x0 = torch.floor(x_dst).long() x1 = x0 + 1 y0 = torch.floor(y_dst).long() y1 = y0 + 1 # 计算每个邻域的权重 wx1 = x_dst - x0.float() wx0 = 1.0 - wx1 wy1 = y_dst - y0.float() wy0 = 1.0 - wy1 # 5. 把源像素的值按权重分配到四个邻域像素 # 先把图像展平为 (B*C, H*W) im_flat = im_t.view(batch_size * channels, src_h * src_w) # 展平坐标和权重 x0_flat = x0.view(-1) x1_flat = x1.view(-1) y0_flat = y0.view(-1) y1_flat = y1.view(-1) wx0_flat = wx0.view(-1) wx1_flat = wx1.view(-1) wy0_flat = wy0.view(-1) wy1_flat = wy1.view(-1) # 计算目标图像的索引(按H*W的顺序) idx_y0x0 = y0_flat * src_w + x0_flat idx_y0x1 = y0_flat * src_w + x1_flat idx_y1x0 = y1_flat * src_w + x0_flat idx_y1x1 = y1_flat * src_w + x1_flat # 初始化目标图像张量 warped_t = torch.zeros(batch_size * channels, src_h * src_w, device=im_t.device) # 用scatter_add_累加权重,这一步是可微分的! warped_t.scatter_add_(1, idx_y0x0.unsqueeze(0).repeat(batch_size*channels, 1), im_flat * wx0_flat * wy0_flat) warped_t.scatter_add_(1, idx_y0x1.unsqueeze(0).repeat(batch_size*channels, 1), im_flat * wx1_flat * wy0_flat) warped_t.scatter_add_(1, idx_y1x0.unsqueeze(0).repeat(batch_size*channels, 1), im_flat * wx0_flat * wy1_flat) warped_t.scatter_add_(1, idx_y1x1.unsqueeze(0).repeat(batch_size*channels, 1), im_flat * wx1_flat * wy1_flat) # 把结果恢复为原来的形状 warped_t = warped_t.view(batch_size, channels, src_h, src_w) # -------------------------- 核心代码结束 -------------------------- # 绘图对比 plt.figure(figsize=(10, 5)) plt.subplot(121) plt.imshow(im_t.data.cpu().numpy()[0,0,:,:], cmap='gray') plt.title("Original Source Image") plt.grid(True) plt.subplot(122) plt.imshow(warped_t.data.cpu().numpy()[0,0,:,:], cmap='gray') plt.title("Forward Warped Image") plt.grid(True) plt.show(block=False) plt.draw()
关键细节说明
- 坐标系统:这里我们用的是像素坐标(x从0到w-1,y从0到h-1),和你原来的归一化坐标不同,但逻辑更直观,适合正向映射。
- 双线性插值:每个源像素映射到目标的浮点坐标后,会把自身的数值按距离权重分配给周围四个整数坐标的像素,这样一个源像素就能影响多个目标像素——如果你的偏移量让源像素映射到离多个目标像素都有距离的位置,自然就能实现非相邻像素的关联。
- 可微分性:
scatter_add_是PyTorch原生支持反向传播的操作,所以整个正向映射过程可以无缝接入神经网络的训练流程,梯度会正确传递回源图像和偏移参数(u、v)。 - 边界处理:这里用
clamp把超出范围的坐标限制在图像内,你也可以根据需求改成丢弃这些像素(比如通过掩码过滤掉超出边界的坐标)。
如果你需要更复杂的映射(比如仿射变换、透视变换),只需要修改x_dst和y_dst的计算逻辑即可,核心的权重分配和scatter累加部分可以复用。
内容的提问来源于stack exchange,提问作者Christopher Grimm
相关产品推荐
相关产品推荐

