如何对BxCxHxW批量图像张量实现高效PyTorch Roll操作?
批量张量Roll操作的高效实现
针对BxCxHxW格式的图像张量和Bx2格式的偏移张量,替代循环的高效实现方案如下:
核心思路
利用PyTorch的高级索引实现批量维度变换,避免逐样本循环带来的性能损耗,完全通过张量运算完成所有样本的Roll操作。
实现代码
import torch def batch_roll(images, shifts, dims=(2, 3)): B, C, H, W = images.shape device = images.device # 生成H、W维度的基础索引 h_idx = torch.arange(H, device=device).unsqueeze(0).repeat(B, 1) # 形状BxH w_idx = torch.arange(W, device=device).unsqueeze(0).repeat(B, 1) # 形状BxW # 对每个样本的索引应用偏移,对齐torch.roll的循环逻辑 dh, dw = shifts[:, 0], shifts[:, 1] h_idx = (h_idx - dh.unsqueeze(1)) % H w_idx = (w_idx - dw.unsqueeze(1)) % W # 扩展索引维度,匹配图像张量的BxCxHxW形状 h_idx = h_idx[:, None, :, None] # 形状Bx1xHx1 w_idx = w_idx[:, None, None, :] # 形状Bx1x1xW # 高级索引完成批量Roll操作 return images[torch.arange(B)[:, None, None, None], :, h_idx, w_idx]
使用示例
# 构造测试数据 batch_size = 4 images = torch.randn(batch_size, 3, 8, 8) # BxCxHxW格式 shifts = torch.tensor([[1, 2], [3, 1], [0, 4], [2, 0]], dtype=torch.long) # Bx2格式的偏移量 # 执行批量Roll rolled_images = batch_roll(images, shifts, dims=(2, 3)) # 验证与循环实现的一致性 for i in range(batch_size): assert torch.allclose(rolled_images[i], images[i].roll(shifts[i].tolist(), (1, 2)))
性能优势
- 完全基于向量化张量运算,消除Python循环的开销,大批次数据下性能提升明显。
- 支持GPU加速运算,充分利用硬件并行能力。
内容的提问来源于stack exchange,提问作者Naufal Suryanto
相关产品推荐
相关产品推荐

