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

求保持相对顺序的打乱累加和序列缺失起始值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的范围,避免遍历所有排列:

  1. 候选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值。
  2. 候选k的验证:
    • 对每个候选k,计算 s_j = b[j] - k,这些值对应原序列中a的累加和;
    • 验证这些s_j能否构成一条完整的累加链:从某个起始s0(满足s0 = a[j])出发,每个后续s都等于前一个s加上对应的a元素,且遍历所有元素。

高效实现代码

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 17:15:54