Python中高效计算将一个向量映射到另一个向量的所有排列
寻找满足
f_arr[perm] = t_arr的有效索引排列(含重复元素场景) 给定两个长度相同的NumPy数组f_arr(源向量)和t_arr(目标向量),需要找出所有索引排列perm,使得f_arr[perm]与t_arr完全相等。向量允许包含重复元素,但无需生成全部排列,现有的代码效率极低,需要基于回溯的高效实现或者优化的Python库来仅生成符合要求的有效排列。
原低效代码示例
import numpy as np import itertools f_arr = np.array([1,2,3,4,3,4], dtype=np.uint8) # 源向量 t_arr = np.array([3,1,4,3,4,2], dtype=np.uint8) # 目标向量 positions = [np.where(f_arr == a)[0] for a in t_arr] for perm in itertools.product(*positions): if len(perm) == len(set(perm)): print(f'{perm} -> {f_arr[list(perm)]}') else: # 仅用于演示无效情况 print(f'非有效排列: {perm}')
原代码运行输出
非有效排列: (2, 0, 3, 2, 3, 1) 非有效排列: (2, 0, 3, 2, 5, 1) 非有效排列: (2, 0, 3, 4, 3, 1) (2, 0, 3, 4, 5, 1) -> [3 1 4 3 4 2] 非有效排列: (2, 0, 5, 2, 3, 1) 非有效排列: (2, 0, 5, 2, 5, 1) (2, 0, 5, 4, 3, 1) -> [3 1 4 3 4 2] 非有效排列: (2, 0, 5, 4, 5, 1) 非有效排列: (4, 0, 3, 2, 3, 1) (4, 0, 3, 2, 5, 1) -> [3 1 4 3 4 2] 非有效排列: (4, 0, 3, 4, 3, 1) 非有效排列: (4, 0, 3, 4, 5, 1) (4, 0, 5, 2, 3, 1) -> [3 1 4 3 4 2] 非有效排列: (4, 0, 5, 2, 5, 1) 非有效排列: (4, 0, 5, 4, 3, 1) 非有效排列: (4, 0, 5, 4, 5, 1)
高效回溯实现
原代码通过itertools.product生成所有可能的索引组合再过滤,会产生大量无效排列(重复索引),效率极低。回溯法可以在生成过程中提前剪枝,只保留有效分支,直接生成符合要求的排列。
import numpy as np from collections import Counter def find_valid_permutations(f_arr, t_arr): # 预处理:统计每个值在f_arr中的索引列表 value_indices = {} for idx, val in enumerate(f_arr): value_indices.setdefault(val, []).append(idx) # 提前校验:t_arr的元素需求是否能被f_arr满足 t_counter = Counter(t_arr) f_counter = Counter(f_arr) for val, cnt in t_counter.items(): if f_counter.get(val, 0) < cnt: return [] # 无有效排列 result = [] used = set() def backtrack(current_idx, current_perm): if current_idx == len(t_arr): result.append(tuple(current_perm)) return target_val = t_arr[current_idx] # 遍历当前目标值对应的未使用索引 for idx in value_indices[target_val]: if idx not in used: used.add(idx) current_perm.append(idx) backtrack(current_idx + 1, current_perm) # 回溯:撤销选择 current_perm.pop() used.remove(idx) backtrack(0, []) return result # 测试示例 f_arr = np.array([1,2,3,4,3,4], dtype=np.uint8) t_arr = np.array([3,1,4,3,4,2], dtype=np.uint8) valid_perms = find_valid_permutations(f_arr, t_arr) for perm in valid_perms: print(f'{perm} -> {f_arr[list(perm)]}')
代码优势
- 提前校验:先检查
t_arr的元素是否在f_arr中有足够数量,直接排除不可能的情况,避免无意义计算。 - 剪枝优化:递归过程中只选择未被使用的索引,不会生成重复索引的无效组合,大幅减少计算量。
- 内存高效:仅存储有效排列,无需生成所有可能组合再过滤。
优化后运行输出
(2, 0, 3, 4, 5, 1) -> [3 1 4 3 4 2] (2, 0, 5, 4, 3, 1) -> [3 1 4 3 4 2] (4, 0, 3, 2, 5, 1) -> [3 1 4 3 4 2] (4, 0, 5, 2, 3, 1) -> [3 1 4 3 4 2]
内容的提问来源于stack exchange,提问作者gabalz
相关产品推荐
相关产品推荐

