如何在PyTorch/Numpy中高效实现类SQL表连接(保留顺序)
问题:在PyTorch或Numpy中高效实现类SQL表连接(保留重复键顺序)
场景说明
初始为未排序、含重复且非连续的键,需保留原始顺序。示例代码如下:
import torch # 初始为未排序、非连续的重复key,需保留顺序 ordered_keys = torch.tensor([1,5,5,3,1,1]).T # 形状6x1 # blackbox()仅接受无重复数组,可排序或未排序 # 输入形状3x1 blackbox_input = ordered_keys.unique() # 输入 = [1, 3, 5] # blackbox()输出对应key的特征 # blackbox_output形状3x2:3个key,每个key含2个特征 blackbox_output = blackbox(blackbox_input) # blackbox_output = [[100,101],[300,301],[500,501]] # 需求:通过PyTorch或Numpy实现类SQL表连接, # 得到形状6x2的输出,与初始顺序一致,值为: # [[100,101], # [500,501], # [500,501], # [300,301], # [100,101], # [100,101]] output = ordered_keys.join(blackbox_output) # <--- 无法运行
实现方案
核心思路是利用唯一值索引映射,通过框架内置的unique方法获取原始键在唯一键数组中的位置,再通过索引直接提取对应特征,全程无循环,效率极高。
PyTorch 实现
import torch # 初始键数组 ordered_keys = torch.tensor([1,5,5,3,1,1]).T # 形状6x1 # 获取唯一键,同时返回原始数组每个元素对应的唯一键索引 unique_keys, inverse_indices = ordered_keys.unique(return_inverse=True) # 调用黑盒函数获取唯一键的特征 blackbox_output = blackbox(unique_keys) # 形状3x2 # 通过索引映射得到最终结果(自动保留原始顺序) output = blackbox_output[inverse_indices]
Numpy 实现
import numpy as np # 初始键数组 ordered_keys = np.array([1,5,5,3,1,1]).reshape(-1,1) # 形状6x1 # 获取唯一键及原始键对应的索引(默认会排序唯一键) unique_keys, inverse_indices = np.unique(ordered_keys, return_inverse=True) # 若需要保留唯一键的首次出现顺序,替换为以下代码: # _, idx = np.unique(ordered_keys, return_index=True) # unique_keys = ordered_keys[np.sort(idx)].squeeze() # inverse_indices = np.searchsorted(unique_keys, ordered_keys.squeeze()) # 调用黑盒函数(根据黑盒输入要求调整数组形状) blackbox_output = blackbox(unique_keys.reshape(-1,1)) # 形状3x2 # 通过索引映射得到结果 output = blackbox_output[inverse_indices]
注意:Numpy的
np.unique默认会对唯一键排序,若需要和PyTorch的unique行为一致(保留键首次出现的顺序),可使用注释中的代码调整。
内容的提问来源于stack exchange,提问作者Sida Zhou
相关产品推荐
相关产品推荐

