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

能否用PyTorch grid_sample实现指定的1D线性插值采样需求?

用PyTorch的grid_sample实现指定的1D线性插值采样

完全可以用grid_sample实现你的需求,核心思路是把1D插值转化为2D插值的特例,并利用PyTorch的广播机制避免重复data。下面是具体实现步骤和代码:

解决思路

  • 适配grid_sample的2D输入要求:给data添加一个虚拟维度,把原本的1D插值问题伪装成2D插值——仅在目标维度(W)做插值,虚拟维度固定坐标不参与插值。
  • 复用广播机制避免重复data:将data的batch维度设为1,让它和index的batch维度B自动广播,无需手动复制data B次。
  • 利用grid_sample的边界处理:通过padding_mode='zeros'直接复用其零填充逻辑,不用自己编写边界判断。

具体代码实现

import torch

# 示例参数与数据
N = 5
W = 10
D = 3
B = 4

data = torch.randn(N, W, D)  # 原始形状: (N, W, D)
index = torch.rand(B) * (W - 1)  # 采样索引,范围可超出[0, W-1],测试边界填充

# 1. 调整data形状适配2D grid_sample输入
# grid_sample的2D输入格式为: (batch_size, channels, height, width)
# 这里把N和D合并为channels,W作为height,添加一个虚拟width维度(大小1),batch设为1
data_2d = data.permute(0, 2, 1).reshape(1, N * D, W, 1)

# 2. 处理索引,归一化到grid_sample要求的[-1, 1]范围
# grid_sample的坐标系统:-1对应输入的最上端/左端,1对应最下端/右端
grid_coords = (index / (W - 1)) * 2 - 1  # 当W=1时需单独处理,此处假设W>1

# 3. 构造grid张量,适配2D grid_sample的格式
# grid形状要求: (batch_size, height, width, 2)
# 这里height=1、width=1,每个grid的坐标为(虚拟维度坐标, 目标插值维度坐标),虚拟维度固定为0(归一化后仍为0)
grid = torch.stack([torch.zeros_like(grid_coords), grid_coords], dim=-1)
grid = grid.view(B, 1, 1, 2)

# 4. 调用grid_sample执行插值
output_2d = torch.nn.functional.grid_sample(
    data_2d,
    grid,
    mode='bilinear',  # 线性插值对应2D的bilinear模式
    padding_mode='zeros',  # 超出范围自动零填充
    align_corners=True  # 确保整数索引对应输入的精确位置
)

# 5. 调整输出形状到目标格式(B, N, D)
output = output_2d.view(B, N, D)

关键细节说明

  • 形状转换逻辑:data从(N, W, D)转为(1, N*D, W, 1),是为了让grid_sample把N个样本的D维特征看作一个整体通道,只在W维度(height方向)做插值。
  • 索引归一化:必须把原始索引范围[0, W-1]转换为[-1, 1],这是grid_sample要求的坐标规范,配合align_corners=True可以保证整数索引对应输入的精确位置。
  • 广播机制:data_2d的batch维度是1,grid的batch维度是B,grid_sample会自动将data_2d广播到B个batch,无需手动复制数据,节省内存。

验证正确性

可以用整数索引验证输出是否符合预期:

# 测试整数索引的情况,应该与data的对应位置完全一致
index_test = torch.tensor([2.0])
grid_coords_test = (index_test / (W - 1)) * 2 - 1
grid_test = torch.stack([torch.zeros_like(grid_coords_test), grid_coords_test], dim=-1).view(1, 1, 1, 2)
output_test = torch.nn.functional.grid_sample(data_2d, grid_test, mode='bilinear', padding_mode='zeros', align_corners=True).view(1, N, D)

print(torch.allclose(output_test[0], data[:, 2, :]))  # 输出True,验证正确

内容的提问来源于stack exchange,提问作者Nagabhushan S N

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 12:42:32