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

如何高效找出数组a等值项对应数组b元素的所有组合?

高效实现同组元素两两组合(无循环,支持CPU/CUDA)

给定已排序数组a和对应数组b,需要提取所有a中值相同的分组内,b元素的所有两两组合。示例如下:

输入:

a = [0, 0, 0, 1, 1, 2, 2, 2, 2]
b = [1, 2, 4, 5, 9, 3, 7, 22, 10]

期望输出:

c = [[1, 2],
     [1, 4],
     [2, 4],
     [5, 9],
     [3, 7],
     [3, 22],
     [3, 10],
     [7, 22],
     [7, 10],
     [22, 10]]

现有基于循环的PyTorch实现:

import torch

a = torch.tensor([0, 0, 0, 1, 1, 2, 2, 2, 2])
b = torch.tensor([1, 2, 4, 5, 9, 3, 7, 22, 10])

jumps = torch.cat((torch.tensor([0]),
                   torch.where(a.diff() > 0)[0] + 1,
                   torch.tensor([len(a)])))
cs = []
for i in range(len(jumps) - 1):
    cs.append(torch.combinations(b[jumps[i]:jumps[i + 1]]))
c = torch.cat(cs)

PyTorch 无循环实现

核心思路是先定位分组边界,生成所有分组内的组合索引对,再通过索引直接从b中提取元素,全程避免显式循环,同时支持CPU和CUDA设备。

import torch

# 自动适配CPU/CUDA设备
device = 'cuda' if torch.cuda.is_available() else 'cpu'
a = torch.tensor([0, 0, 0, 1, 1, 2, 2, 2, 2], device=device)
b = torch.tensor([1, 2, 4, 5, 9, 3, 7, 22, 10], device=device)

# 1. 计算每个分组的起始位置和长度
diff = a.diff()
group_starts = torch.cat((torch.tensor([0], device=device), torch.where(diff > 0)[0] + 1))
group_lengths = torch.diff(torch.cat((group_starts, torch.tensor([len(a)], device=device))))

# 2. 生成全局索引对:先为每个分组生成相对组合索引,再加上分组起始偏移
offsets = group_starts.repeat_interleave(group_lengths * (group_lengths - 1) // 2)
rel_indices = torch.cat([torch.combinations(torch.arange(l, device=device)) for l in group_lengths])
global_indices = offsets[:, None] + rel_indices

# 3. 从b中提取对应元素得到最终结果
c = b[global_indices]

复杂度说明

该方案时间复杂度为O(m²),其中m是a中最大等值分组的长度。仅针对每个分组内的元素生成组合,不会触发全数组级别的O(n²)计算,性能更高效。


NumPy 无循环实现

利用NumPy的向量化操作完成分组划分与组合生成,逻辑与PyTorch方案一致:

import numpy as np

a = np.array([0, 0, 0, 1, 1, 2, 2, 2, 2])
b = np.array([1, 2, 4, 5, 9, 3, 7, 22, 10])

# 1. 计算分组起始位置和长度
diff = np.diff(a)
group_starts = np.concatenate(([0], np.where(diff > 0)[0] + 1))
group_lengths = np.diff(np.concatenate((group_starts, [len(a)])))

# 2. 生成全局索引对:先生成分组内的相对索引,筛选i<j的组合后加上分组偏移
offsets = np.repeat(group_starts, group_lengths * (group_lengths - 1) // 2)
rel_indices = np.concatenate([np.array(np.meshgrid(np.arange(l), np.arange(l))).T.reshape(-1,2) for l in group_lengths])
rel_indices = rel_indices[rel_indices[:, 0] < rel_indices[:, 1]]
global_indices = offsets[:, None] + rel_indices

# 3. 提取结果
c = b[global_indices]

内容的提问来源于stack exchange,提问作者Julius

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 12:02:16