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

如何利用索引图像对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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 14:47:46