求膨胀k近邻算法的PyTorch广播优化方案以提升执行速度
优化Dilated K-Nearest Neighbors的嵌套循环实现
核心思路
原代码通过三层嵌套循环逐元素筛选符合索引 % dilation == k % dilation的近邻,效率极低。我们可以利用PyTorch的广播机制和张量操作实现批量处理,完全移除Python级循环,充分利用硬件加速。
优化后的代码
import torch dilation = 3 nbd_size = 5 # 修正原randint参数格式(原写法不符合PyTorch规范) knn_key = torch.randint(0, 30, (64, 12, 198, 100), device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')) # 1. 生成k维度的模dilation值,广播到与knn_key匹配的形状 k_coords = torch.arange(knn_key.shape[2], device=knn_key.device) k_mod = k_coords % dilation # 扩展为(1,1,198,1),和knn_key的(64,12,198,100)做元素级广播对比 k_mod = k_mod.view(1, 1, -1, 1) # 2. 计算所有符合条件的掩码 mask = (knn_key % dilation) == k_mod # 3. 给不符合条件的位置赋予极大索引值,确保topk只选原顺序的符合项 l_idx = torch.arange(knn_key.shape[3], device=knn_key.device).expand_as(knn_key) l_idx[~mask] = knn_key.shape[3] # 设为超过l维度长度的值,不会被topk选中 # 4. 提取每个(i,j,k)的前nbd_size个符合条件的近邻索引 # largest=False取最小的索引,对应原l维度的先后顺序 top_l_indices = torch.topk(l_idx, nbd_size, dim=-1, largest=False)[1] dilated_keys = torch.gather(knn_key, dim=-1, index=top_l_indices)
关键细节说明
- 广播匹配:将k维度的模值扩展为
(1,1,198,1),和knn_key的四维形状广播对齐,实现一次性完成所有元素的条件判断,替代循环中的逐元素检查。 - 保留原顺序:通过给不符合条件的位置赋极大值,
topk选取最小的nbd_size个索引时,只会保留原l维度顺序中最先出现的符合条件的元素,完全匹配原代码逻辑。 - 性能提升:张量操作由PyTorch底层优化(支持CUDA),相比嵌套循环,速度可提升数十倍甚至上百倍,尤其在大张量场景下效果显著。
注:原代码中torch.randint的参数写法错误,已修正为PyTorch标准格式torch.randint(low, high, shape)。
内容的提问来源于stack exchange,提问作者Aleph
相关产品推荐
相关产品推荐

