求保持相对顺序的打乱累加和序列缺失起始值k的高效算法
高效求解k值的算法问题
问题背景
给定两个长度为n的float数组a和b,满足以下条件:
- a数组可包含正负数值;
- 原有序列中,b是a的累加和(cumsum);
- 实际b的第一个元素满足
b[0] = a[0] + k(且b[0]≠a[0]); - a和b被随机打乱,但两者的相对顺序保持一致:即原a中第i个元素如果被放到a的第j位,原b中第i个元素也会被放到b的第j位。
需要找到高效算法计算k值。
现有朴素实现的问题
目前的朴素实现通过遍历所有排列验证,但当n≥10时运行极慢,代码如下:
import numpy as np import itertools def get_starting_point(a, b): for msk in itertools.permutations(range(len(a))): # NOTE: n≥10时运行极慢 new_a = a[list(msk)] new_b = b[list(msk)] k = new_b[0] - new_a[0] new_a = np.cumsum(new_a) + k if np.nansum(np.abs(new_b - new_a)) < 0.001: return k return None
高效解法思路
我们可以利用问题的数学性质大幅缩小候选k的范围,避免遍历所有排列:
- 候选k的筛选:
- 原序列最后一个b元素满足
b[-1] = k + sum(a),因此k可表示为b[i] - sum(a)(i为b中任意位置); - 原序列第一个b元素满足
b[0] = k + a[0],因此k也可表示为b[i] - a[i](i为任意对应位置); - 取这两个候选集合的交集,即可得到数量极少的候选k值。
- 原序列最后一个b元素满足
- 候选k的验证:
- 对每个候选k,计算
s_j = b[j] - k,这些值对应原序列中a的累加和; - 验证这些
s_j能否构成一条完整的累加链:从某个起始s0(满足s0 = a[j])出发,每个后续s都等于前一个s加上对应的a元素,且遍历所有元素。
- 对每个候选k,计算
高效实现代码
import numpy as np def get_k(a, b): n = len(a) if n == 0: return None sum_a = np.sum(a) # 生成候选k集合 candidates1 = set(round(b[j] - sum_a, 2) for j in range(n)) candidates2 = set(round(b[j] - a[j], 2) for j in range(n)) candidates = candidates1.intersection(candidates2) # 验证每个候选k for k in candidates: s_list = [round(b[j] - k, 2) for j in range(n)] # 构建s到a的映射,处理重复值 s_to_a = {} for s, aj in zip(s_list, a): s_rounded = round(s, 2) aj_rounded = round(aj, 2) if s_rounded not in s_to_a: s_to_a[s_rounded] = [] s_to_a[s_rounded].append(aj_rounded) # 找所有可能的起始点:s0对应的a元素等于s0本身 start_points = [] for s in s_to_a: if any(abs(aj - s) < 1e-6 for aj in s_to_a[s]): start_points.append(s) if not start_points: continue # 验证每个起始点的累加链 for s0 in start_points: used = {s0} temp_map = {key: val.copy() for key, val in s_to_a.items()} # 移除起始点对应的a元素 temp_map[s0].remove(s0) if not temp_map[s0]: del temp_map[s0] count = 1 valid = True while count < n: found = None # 寻找符合条件的下一个s for s in list(temp_map.keys()): for aj in temp_map[s]: prev_s = round(s - aj, 2) if prev_s in used: found = (s, aj) break if found: break if not found: valid = False break # 更新已使用集合和映射 s_found, aj_found = found used.add(s_found) temp_map[s_found].remove(aj_found) if not temp_map[s_found]: del temp_map[s_found] count += 1 if valid: return round(k, 2) # 若交集候选无结果,遍历所有b[j]-a[j]候选 for k in candidates2: k_rounded = round(k, 2) s_list = [round(b[j] - k_rounded, 2) for j in range(n)] s_to_a = {} for s, aj in zip(s_list, a): s_rounded = round(s, 2) aj_rounded = round(aj, 2) if s_rounded not in s_to_a: s_to_a[s_rounded] = [] s_to_a[s_rounded].append(aj_rounded) start_points = [] for s in s_to_a: if any(abs(aj - s) < 1e-6 for aj in s_to_a[s]): start_points.append(s) if not start_points: continue for s0 in start_points: used = {s0} temp_map = {key: val.copy() for key, val in s_to_a.items()} temp_map[s0].remove(s0) if not temp_map[s0]: del temp_map[s0] count = 1 valid = True while count < n: found = None for s in list(temp_map.keys()): for aj in temp_map[s]: prev_s = round(s - aj, 2) if prev_s in used: found = (s, aj) break if found: break if not found: valid = False break s_found, aj_found = found used.add(s_found) temp_map[s_found].remove(aj_found) if not temp_map[s_found]: del temp_map[s_found] count += 1 if valid: return k_rounded return None
测试验证
使用提供的生成函数测试算法正确性:
def get_a_b_k(n=14): a = np.round(np.random.uniform(low=-10, high=10, size=(n,)), 2) b = np.cumsum(a) prob = np.random.uniform(0,1) if prob < 0.4: k = np.round(np.random.uniform(-10,10), 2) elif prob < 0.6: # k same as the last b. k = b[n-1] a[n-2] -= k else: # k same as one of b's idx = np.random.choice(n, size=1) k = b[idx] a[idx] -= k b = np.cumsum(a) msk = np.random.choice(n, size=n, replace=False) # Randomly generated mask of size n. return a[msk], b[msk] + k, k # 测试10次 for _ in range(10): a, b, expected_k = get_a_b_k(n=14) computed_k = get_k(a, b) print(f"Expected k: {expected_k}, Computed k: {computed_k}, Match: {abs(expected_k - computed_k) < 1e-6}")
内容的提问来源于stack exchange,提问作者Gerry
相关产品推荐
相关产品推荐

