You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

神经网络框架中可微分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()

关键细节说明

  1. 坐标系统:这里我们用的是像素坐标(x从0到w-1,y从0到h-1),和你原来的归一化坐标不同,但逻辑更直观,适合正向映射。
  2. 双线性插值:每个源像素映射到目标的浮点坐标后,会把自身的数值按距离权重分配给周围四个整数坐标的像素,这样一个源像素就能影响多个目标像素——如果你的偏移量让源像素映射到离多个目标像素都有距离的位置,自然就能实现非相邻像素的关联。
  3. 可微分性:scatter_add_是PyTorch原生支持反向传播的操作,所以整个正向映射过程可以无缝接入神经网络的训练流程,梯度会正确传递回源图像和偏移参数(u、v)。
  4. 边界处理:这里用clamp把超出范围的坐标限制在图像内,你也可以根据需求改成丢弃这些像素(比如通过掩码过滤掉超出边界的坐标)。

如果你需要更复杂的映射(比如仿射变换、透视变换),只需要修改x_dst和y_dst的计算逻辑即可,核心的权重分配和scatter累加部分可以复用。

内容的提问来源于stack exchange,提问作者Christopher Grimm

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.14 07:23:58