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

如何在PyTorch中按指定像素坐标偏移规则实现图像变换

PyTorch实现可微空间偏移采样的方法

核心原理

你要实现的操作本质是基于预测偏移场的图像重采样,PyTorch内置的torch.nn.functional.grid_sample接口原生支持GPU加速和自动求导,可以直接集成到训练流水线中,不需要自定义算子。

前置说明

  • 所有输入张量默认格式为PyTorch标准的[batch_size, channels, height, width](即NCHW格式)
  • 假设X、Y是sigmoid激活后的输出,取值范围为[0,1],符合你给出的X-0.5偏移计算逻辑
  • 采样默认使用双线性插值保证可微性,超出图像边界的坐标自动用边缘像素填充,你可以按需调整参数

实现代码

import torch
import torch.nn.functional as F

def offset_sample(P: torch.Tensor, X: torch.Tensor, Y: torch.Tensor, scale: float) -> torch.Tensor:
    """
    按照给定偏移规则对P进行重采样得到P'
    参数:
        P: 输入图像,形状为[N, 1, H, W]
        X: 水平偏移场,形状为[N, 1, H, W]
        Y: 垂直偏移场,形状为[N, 1, H, W]
        scale: 偏移缩放系数
    返回:
        P_prime: 重采样后的图像,形状为[N, 1, H, W]
    """
    N, _, H, W = P.shape
    device = P.device

    # 生成原始像素坐标网格
    y_grid, x_grid = torch.meshgrid(
        torch.arange(H, device=device),
        torch.arange(W, device=device),
        indexing='ij'
    )
    # 扩展为batch维度 形状变为[N, H, W]
    x_grid = x_grid.unsqueeze(0).repeat(N, 1, 1).float()
    y_grid = y_grid.unsqueeze(0).repeat(N, 1, 1).float()

    # 计算偏移后的坐标,去掉X/Y的通道维度
    offset_x = scale * (X.squeeze(1) - 0.5)
    offset_y = scale * (Y.squeeze(1) - 0.5)
    x_new = x_grid + offset_x
    y_new = y_grid + offset_y

    # 将坐标归一化到grid_sample要求的[-1, 1]区间
    x_norm = (x_new / (W - 1)) * 2 - 1
    y_norm = (y_new / (H - 1)) * 2 - 1

    # 拼接为采样网格 形状为[N, H, W, 2],最后一维顺序是(x,y)
    grid = torch.stack([x_norm, y_norm], dim=-1)

    # 执行采样,align_corners设置为True保证像素坐标对齐
    P_prime = F.grid_sample(
        P,
        grid,
        mode='bilinear',
        padding_mode='border',
        align_corners=True
    )
    return P_prime

使用示例

# 示例参数
batch_size = 2
H, W = 256, 256
scale = 10.0

# 随机生成模拟输入,可直接放到GPU上运行
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
P = torch.randn(batch_size, 1, H, W).to(device)
X = torch.sigmoid(torch.randn(batch_size, 1, H, W)).to(device)
Y = torch.sigmoid(torch.randn(batch_size, 1, H, W)).to(device)

# 调用函数
P_prime = offset_sample(P, X, Y, scale)

# 可直接接入损失函数计算,梯度自动回传
loss = P_prime.mean()
loss.backward()

可调整参数说明

  • 若需要调整采样精度,可将mode参数改为'nearest'(最近邻插值,不可微)或'bicubic'(双三次插值)
  • 若需要处理超出边界的坐标,可将padding_mode改为'zero'(补零)或'reflection'(镜像填充)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 06:06:03