如何高效根据给定坐标批量提取多维张量的对应数据?
5维张量批量坐标高效索引方案
原写法性能瓶颈
逐坐标循环取值的实现中,所有索引操作都在Python层逐个触发,框架无法对计算做批量调度,大量细碎操作的调度开销会拉低整体运行速度,坐标点数量越多性能损耗越明显。
零循环高效实现
直接将所有坐标拼接为批量索引张量,利用框架原生高级索引能力一次性完成所有取值,操作全部下沉到底层C/CUDA核执行,无额外Python循环开销,以下是PyTorch框架下的实现(3D视觉领域该类dense volume索引场景基本均基于PyTorch实现):
import torch # 1. 将所有独立坐标堆叠为形状为 [N, 4] 的批量索引矩阵,N为坐标点总数 coords_batch = torch.stack(list(coords), dim=0) # 2. 一次性完成批量索引,结果形状为 [N, C],C为张量通道维长度 # 索引维度对应关系和原写法完全一致:第0维取coord[3]、第2维取coord[0]、第3维取coord[1]、第4维取coord[2],通道维全取 feats = dense_volume1[ coords_batch[:, 3], # 对应原张量第0维的索引 :, # 通道维全部保留 coords_batch[:, 0], # 对应原张量第2维的索引 coords_batch[:, 1], # 对应原张量第3维的索引 coords_batch[:, 2] # 对应原张量第4维的索引 ]
效果说明
- 计算结果和原循环写法逐点append后拼接的结果完全数值对齐,无精度差异
- 性能通常比原循环实现快1~2个数量级,坐标点规模越大提速越明显
- 同时适配CPU、GPU张量,不需要额外修改设备相关逻辑
内容的提问来源于stack exchange,提问作者andiska
相关产品推荐
相关产品推荐

