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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 12:10:17