如何在PyTorch中无需循环从Tensor指定位置批量提取图像Patch
高效实现方案
有两种纯PyTorch原生的无循环实现方式,其中高级索引广播法性能最优,不需要引入额外算子:
方法1:高级索引广播实现
核心思路是预先批量生成所有Patch对应的像素坐标,利用PyTorch的广播机制一次性完成索引,完全避免Python层循环。
完整示例代码如下:
import torch # 原始输入 d = torch.rand(4, 64, 64) xy = torch.tensor([(15, 21), (30, 59), (40, 5), (20, 25)]) x, y = xy.T m = 1 patch_size = 2 * m + 1 # 生成行列方向的偏移量 row_offset = torch.arange(-m, m+1, device=d.device) col_offset = torch.arange(-m, m+1, device=d.device) # 广播生成所有Patch对应的全局坐标 batch_idx = torch.arange(d.shape[0], device=d.device)[:, None, None] # 形状 [N, 1, 1] row_idx = x[:, None, None] + row_offset[None, :, None] # 形状 [N, patch_size, 1] col_idx = y[:, None, None] + col_offset[None, None, :] # 形状 [N, 1, patch_size] # 一次性索引获取所有Patch o = d[batch_idx, row_idx, col_idx] print(o.size()) # 输出 torch.Size([4, 3, 3])
注意:如果坐标存在边缘越界的情况,可以先通过
torch.nn.functional.pad对输入张量的后两维填充m个像素,再将所有x、y坐标统一加m后再做索引,即可避免越界错误。
方法2:unfold滑窗实现(适合多Patch提取场景)
如果需要对单张图提取多个Patch,可以先用torch.nn.functional.unfold把所有可能的Patch提前展开,再根据坐标索引对应Patch:
from torch.nn import functional as F # 滑动窗口展开所有3x3 Patch,展开后形状为 [N, 3*3, 有效滑窗数量] unfolded = F.unfold(d.unsqueeze(1), kernel_size=patch_size, padding=0) # 计算坐标对应的Patch在展开后维度的位置 pos = (x - m) * (d.shape[2] - 2*m) + (y - m) # 索引后调整形状 o = unfolded[torch.arange(d.shape[0]), :, pos].reshape(-1, patch_size, patch_size)
这种方法更适合单图多Patch的场景,单图单Patch场景下性能略逊于方法1。
两种方法的输出都和原有循环实现的结果完全一致,在批量较大的场景下速度会比Python循环快几十到上百倍。
内容的提问来源于stack exchange,提问作者ravi
相关产品推荐
相关产品推荐

