如何利用索引图像对PyTorch张量批量图像进行索引?
问题:根据索引图像从批量图像中提取像素
我有一个形状为(B, W, H)的PyTorch张量图像批量M,还有一个尺寸为(W, H)、像素为索引值的图像I。需要生成一个(W, H)的图像,每个像素取自M中对应索引的图像(遵循I的索引规则)。
示例
给定形状为(3, 4, 8)的M:
tensor([[[ 0., 0., 0., 0., 0., 0., 0., 0.], [ 0., 0., 0., 0., 0., 0., 0., 0.], [ 0., 0., 0., 0., 0., 0., 0., 0.], [ 0., 0., 0., 0., 0., 0., 0., 0.]], [[-1., -1., -1., -1., -1., -1., -1., -1.], [-1., -1., -1., -1., -1., -1., -1., -1.], [-1., -1., -1., -1., -1., -1., -1., -1.], [-1., -1., -1., -1., -1., -1., -1., -1.]], [[-2., -2., -2., -2., -2., -2., -2., -2.], [-2., -2., -2., -2., -2., -2., -2., -2.], [-2., -2., -2., -2., -2., -2., -2., -2.], [-2., -2., -2., -2., -2., -2., -2., -2.]]])
以及形状为(4, 8)的I:
tensor([[2, 0, 2, 0, 1, 0, 1, 0], [2, 2, 1, 0, 0, 2, 1, 0], [2, 0, 0, 2, 1, 1, 0, 0], [0, 1, 0, 0, 2, 0, 2, 1]], dtype=torch.int32)
得到的结果图像为:
tensor([[-2., 0., -2., 0., -1., 0., -1., 0.], [-2., -2., -1., 0., 0., -2., -1., 0.], [-2., 0., 0., -2., -1., -1., 0., 0.], [ 0., -1., 0., 0., -2., 0., -2., -1.]])
PyTorch实现方案
当M为(B, W, H)格式时
利用PyTorch的高级索引直接实现,通过生成空间维度的索引矩阵,与I组合完成像素提取:
import torch # 假设已定义M和I W, H = I.shape # 生成W和H维度的全索引矩阵 w_idx = torch.arange(W).unsqueeze(1).repeat(1, H) h_idx = torch.arange(H).unsqueeze(0).repeat(W, 1) # 按索引提取像素 result = M[I, w_idx, h_idx]
I作为批量维度的索引,w_idx和h_idx对应每个空间位置的坐标,三者组合后即可精准提取对应像素。
当M为(W, H, B)格式时
这种维度顺序下,使用torch.gather可以更简洁地完成任务:
# 将M转换为(W, H, B)格式 M_transposed = M.permute(1, 2, 0) # 扩展I的维度以匹配M_transposed,在B维度上按索引取值 result = torch.gather(M_transposed, dim=2, index=I.unsqueeze(-1)).squeeze(-1)
I.unsqueeze(-1)是为了让索引维度与目标张量对齐,gather会在指定维度(这里是B维度)上按照索引提取元素,最后通过squeeze去掉多余的维度得到(W, H)结果。
NumPy实现方案
当M为(B, W, H)格式时
利用NumPy的花式索引,逻辑与PyTorch版本一致:
import numpy as np # 假设已将M和I转换为NumPy数组M_np、I_np W, H = I_np.shape # 生成空间维度索引矩阵 w_idx = np.arange(W)[:, np.newaxis].repeat(H, axis=1) h_idx = np.arange(H)[np.newaxis, :].repeat(W, axis=0) # 提取像素 result_np = M_np[I_np, w_idx, h_idx]
当M为(W, H, B)格式时
使用np.take_along_axis函数完成索引提取:
# 将M转换为(W, H, B)格式 M_np_transposed = np.transpose(M_np, (1, 2, 0)) # 扩展索引维度以匹配目标张量 I_np_expanded = I_np[..., np.newaxis] # 按索引提取元素并压缩维度 result_np = np.take_along_axis(M_np_transposed, I_np_expanded, axis=2).squeeze(axis=2)
内容的提问来源于stack exchange,提问作者arthur.sw
相关产品推荐
相关产品推荐

