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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 22:30:51