如何无需循环实现张量(图像)的差异化位移操作?
批量实现张量图像的个性化循环位移
要批量给每张图像施加不同的循环位移(无需遍历单张图像),可以用以下两种高效方法:
方法1:用torch.vmap批量映射单样本操作
torch.vmap是PyTorch的批量映射工具,能把针对单张图像的位移逻辑自动扩展到整个批量,代码简洁且效率高(适合PyTorch 1.10及以上版本)。
示例代码
import torch from torch import vmap # 构造示例批量图像张量 img_batch = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8]).view(2, 2, 2) # 定义每个样本的位移:格式为(y轴位移量, x轴位移量) shifts = torch.tensor([[0, 1], [1, 0]]) # 定义单张图像的循环位移函数 def roll_single(img, shift): # img: 单张图像,形状[H, W];shift: 该图像的位移参数 return torch.roll(img, shifts=shift.tolist(), dims=(0, 1)) # 批量执行位移 result = vmap(roll_single)(img_batch, shifts) print(result)
输出结果
tensor([[[2, 1], [4, 3]], [[7, 8], [5, 6]]])
方法2:用高级索引手动构造位移索引
如果你的PyTorch版本较低,不支持vmap,可以通过构造每个样本的位移索引矩阵,结合torch.gather实现批量位移。
示例代码
import torch img_batch = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8]).view(2, 2, 2) B, H, W = img_batch.shape # 定义每个样本的x轴、y轴位移量 x_shifts = torch.tensor([1, 0]) y_shifts = torch.tensor([0, 1]) # 生成y轴位移后的索引矩阵 y_indices = torch.arange(H, device=img_batch.device).repeat(B, 1) y_indices = (y_indices - y_shifts[:, None]) % H # 取模实现循环位移 # 扩展维度以匹配gather的输入要求 y_indices = y_indices[:, :, None].repeat(1, 1, W) # 先完成y轴方向的位移 img_y_rolled = torch.gather(img_batch, dim=1, index=y_indices) # 生成x轴位移后的索引矩阵 x_indices = torch.arange(W, device=img_batch.device).repeat(B, 1) x_indices = (x_indices - x_shifts[:, None]) % W x_indices = x_indices[:, None, :].repeat(1, H, 1) # 完成x轴方向的位移 final_result = torch.gather(img_y_rolled, dim=2, index=x_indices) print(final_result)
输出结果
和方法1完全一致。
方法对比
vmap方法:代码简洁易读,底层自动优化批量逻辑,性能最优,推荐优先使用。- 高级索引方法:兼容性强,无需依赖高版本PyTorch,但代码相对繁琐。
内容的提问来源于stack exchange,提问作者Erney Ramírez
相关产品推荐
相关产品推荐

