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

如何无需循环实现张量(图像)的差异化位移操作?

批量实现张量图像的个性化循环位移

要批量给每张图像施加不同的循环位移(无需遍历单张图像),可以用以下两种高效方法:

方法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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 05:53:32