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

求膨胀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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 22:12:04