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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 21:00:03