PyTorch中按图像批量提取指定像素的最优实现方法
PyTorch批量按自定义索引提取图像像素的最优方案
问题背景
我有一批形状为(n_images, height, width, channels)的图像张量,以及一批对应每张图像的(x,y)索引张量(形状为(n_images, n_samples, 2),每张图像的索引各不相同),无法通过简单索引直接提取目标像素。需要获取形状为(n_images, n_samples, channels)的结果张量,包含每张图像对应索引位置的像素颜色。
示例代码:
import torch n_images = 4 width = 100 height = 100 channels = 3 n_samples = 30 images = torch.rand((n_images, height, width, channels)) indices = (torch.rand((n_images, n_samples, 2)) * width).to(torch.int32) # 期望效果:result.shape = (n_images, n_samples, 3) # result = images[indices] # 此写法无法直接生效
当前实现方案
我已经实现了一种可行方案,但希望用更通用的PyTorch内置函数来替代:
# 当前实现,但希望改用通用Torch函数 xs = indices.reshape((-1, 2))[:, 0] ys = indices.reshape((-1, 2))[:, 1] ix = torch.arange(n_images, dtype=torch.int32) ix = ix[..., None].expand((-1, n_samples)).flatten() result = images[ix, ys, xs].reshape((n_images, n_samples, 3))
最优/更简洁的实现方式
你的当前方案本质上已经是高效的高级索引实现,不过可以简化写法,同时利用PyTorch的广播机制让代码更简洁:
方法1:简化高级索引写法
# 拆分索引的y(对应height维度)和x(对应width维度) ys = indices[..., 1] xs = indices[..., 0] # 生成批量索引,利用广播自动匹配n_samples维度 batch_idx = torch.arange(n_images, device=images.device)[:, None] # 直接通过高级索引提取,自动对齐维度 result = images[batch_idx, ys, xs]
此写法和你的原方案效率完全一致,但无需额外的reshape/flatten操作,代码更简洁直观,同时保留了高级索引的硬件加速优势。
方法2:使用torch.gather(通用维度索引场景)
如果需要适配更通用的维度动态索引场景,可以借助torch.gather,但需要调整张量维度以匹配函数要求:
# 将图像张量转置为 (n_images, channels, height, width),方便按通道维度聚合 images_transposed = images.permute(0, 3, 1, 2) # 扩展索引维度以匹配转置后的张量结构 indices_expanded = indices.unsqueeze(1) # 先在height维度提取,再在width维度提取 gathered_height = torch.gather(images_transposed, dim=2, index=indices_expanded[..., 1].unsqueeze(-1).expand(-1, channels, -1, 1)) result = torch.gather(gathered_height, dim=3, index=indices_expanded[..., 0].unsqueeze(-1).expand(-1, channels, -1, 1)) # 调整回目标形状 result = result.squeeze(-1).permute(0, 2, 1)
注意:这种写法效率略低于直接高级索引,因为涉及多次转置和维度扩展,更适合需要动态指定提取维度的通用场景。
性能说明
直接高级索引的方式(方法1)是PyTorch处理此类场景的最优实现,底层会利用CUDA/CPU的硬件加速,执行效率最高。你的原方案和方法1的性能几乎无差异,只是写法更简洁。
内容的提问来源于stack exchange,提问作者Flooo
相关产品推荐
相关产品推荐

