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

如何仅用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}")

关键细节说明

  1. 坐标转换:严格对应PyTorch的align_corners参数逻辑,确保坐标映射完全一致。
  2. 批量处理:通过batch_idx实现批量维度的索引广播,避免循环处理每个样本。
  3. 边界处理:用np.clip限制坐标范围,实现constant和border填充模式,超出图像范围的像素会被截断到边界(constant模式下超出部分最终权重为0,等价于填充0)。
  4. 插值逻辑:双线性插值的权重计算完全遵循数学公式,保证结果和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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 14:42:18