如何高效找出数组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
相关产品推荐
相关产品推荐

