如何用PyTorch实现各通道指定像素读取并写入目标图像?
如何用PyTorch实现张量跨通道指定位置的像素读写?
我有一个形状为[C, H, W]的张量,其中C为通道数(不一定是3)。另有两组索引数组,希望从源图像A的对应通道指定索引位置读取像素,再写入目标图像B的对应通道指定索引位置。例如:源图像A与目标图像B,C=3,H、W=64,indicesRead = [[0, 0], [13, 15], [32, 43]],indicesWrite = [[7, 5], [1, 1], [4, 4]],即读取A的通道0的[0,0]像素、通道1的[13,15]像素、通道2的[32,43]像素,分别写入B的通道0的[7,5]、通道1的[1,1]、通道2的[4,4]位置。请问能否用PyTorch方法实现,还是必须手动遍历张量?
当然可以用PyTorch的高级索引实现,完全不需要手动遍历。核心思路是把通道索引和对应的位置索引组合成符合PyTorch索引规则的张量,一次性完成读写操作。
实现步骤
- 将索引数组转换为PyTorch张量,拆分出高度(h)和宽度(w)维度的索引
- 生成与通道数匹配的通道索引张量
- 利用高级索引批量读取源张量A的指定像素
- 再通过高级索引将像素批量写入目标张量B的指定位置
示例代码
import torch # 模拟输入:C=3, H=64, W=64 C, H, W = 3, 64, 64 A = torch.randn(C, H, W) # 源张量 B = torch.zeros(C, H, W) # 目标张量 # 给定的索引数组 indicesRead = [[0, 0], [13, 15], [32, 43]] indicesWrite = [[7, 5], [1, 1], [4, 4]] # 转换索引为PyTorch张量 indicesRead = torch.tensor(indicesRead) indicesWrite = torch.tensor(indicesWrite) # 拆分h和w方向的索引 read_h, read_w = indicesRead[:, 0], indicesRead[:, 1] write_h, write_w = indicesWrite[:, 0], indicesWrite[:, 1] # 生成通道索引:每个通道对应一个索引,形状为[C] channel_idx = torch.arange(C) # 批量读取A中指定位置的像素 pixels = A[channel_idx, read_h, read_w] # 批量写入B的指定位置 B[channel_idx, write_h, write_w] = pixels # 验证结果(可选) assert torch.allclose(A[0, 0, 0], B[0, 7, 5]) assert torch.allclose(A[1, 13, 15], B[1, 1, 1]) assert torch.allclose(A[2, 32, 43], B[2, 4, 4])
关键说明
- PyTorch高级索引支持多维度索引组合,只要各维度索引形状匹配即可。这里
channel_idx、read_h、read_w都是长度为C的一维张量,组合后能精准定位每个通道的指定位置。 - 这种向量化操作的效率远高于手动遍历,通道数C越大,优势越明显。
- 如果索引是numpy数组,只需用
torch.from_numpy()转换即可,无需额外处理。
内容的提问来源于stack exchange,提问作者Martin Perry
相关产品推荐
相关产品推荐

