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

求O(n polylog n)复杂度的连续子数组计数算法实现

优化连续子数组和统计算法至O(n log n)复杂度

问题背景

我们需要实现一个算法,统计数组中元素和处于[tmin, tmax]范围内的非空连续子数组数量。程序要求从input.txt读取输入,将结果写入output.txt。当前实现为O(n²)复杂度,需优化至O(n log n)级别。尝试过二分查找、AVL树但未成功,特寻求解决方案。

现有O(n²)实现代码

import time

def number_of_allowable_intervals(input_file_path, output_file_path):
    # Open input file and read values
    with open(input_file_path, 'r') as input_file:
        # Read the size of the array and allowable range
        n, tmin, tmax = map(int, input_file.readline().strip().split(','))
        A = list(map(int, input_file.readline().replace(',', '').split()))

    # Count all subarrays within the range [tmin, tmax]
    sub_array_count = 0

    start_time = time.time()

    for i in range(len(A)):
        current_sum = 0
        for j in range(i, len(A)):
            current_sum += A[j]

            # Check if the sum is within the range [tmin, tmax]
            if tmin <= current_sum <= tmax:
                sub_array_count += 1

    end_time = time.time()

    print("Runtime:", end_time - start_time)

    with open(output_file_path, 'w') as output_file:
        output_file.write(str(sub_array_count))

if __name__ == "__main__":
    input_path = "input.txt"
    output_path = "output.txt"
    number_of_allowable_intervals(input_path, output_path)

输入样例与预期结果

input.txt内容:

10, 0, 0
-1, 1, -1, 1, -1, 1, -1, 1, -1, 1000

预期输出:20

优化方案:分治法实现O(n log²n)复杂度

分治法核心是将数组递归分割为左右两部分,分别统计左右内部的符合条件子数组,再统计跨越中间点的子数组数量,三者相加得到总计数。跨中间的子数组统计通过排序前缀和+二分查找实现高效计算,整体满足O(n polylog n)要求。

具体步骤:

  1. 递归分割:将数组从中间分为左右两半,递归计算左右各自的符合条件子数组数量。
  2. 统计跨中间的子数组:
    • 计算左半部分从中间位置向左的所有前缀和(包含中间点)。
    • 计算右半部分从中间位置向右的所有前缀和(包含中间点的下一个位置)。
    • 将左半前缀和排序,对右半每个前缀和s_r,用二分查找统计左半中满足 tmin - s_r ≤ s_l ≤ tmax - s_r 的s_l数量。
  3. 合并结果:将左右内部计数与跨中间计数相加,得到总结果。

优化后的代码实现

import time
import bisect

def count_cross_subarrays(arr, left, mid, right, tmin, tmax):
    # 统计左半部分从mid向左的前缀和(包含mid)
    left_prefix = []
    current_sum = 0
    for i in range(mid, left-1, -1):
        current_sum += arr[i]
        left_prefix.append(current_sum)
    # 排序左半前缀和,用于二分查找
    left_prefix.sort()

    # 统计右半部分从mid+1向右的前缀和(包含mid+1),并计算符合条件的数量
    cross_count = 0
    current_sum = 0
    for i in range(mid+1, right+1):
        current_sum += arr[i]
        # 需要找到 left_prefix 中满足 tmin - current_sum <= s <= tmax - current_sum 的数量
        lower = tmin - current_sum
        upper = tmax - current_sum
        # bisect_left找第一个>=lower的索引,bisect_right找第一个>upper的索引,差值即为数量
        left_idx = bisect.bisect_left(left_prefix, lower)
        right_idx = bisect.bisect_right(left_prefix, upper)
        cross_count += right_idx - left_idx
    return cross_count

def count_subarrays_recursive(arr, left, right, tmin, tmax):
    if left == right:
        return 1 if tmin <= arr[left] <= tmax else 0
    
    mid = (left + right) // 2
    left_count = count_subarrays_recursive(arr, left, mid, tmin, tmax)
    right_count = count_subarrays_recursive(arr, mid+1, right, tmin, tmax)
    cross_count = count_cross_subarrays(arr, left, mid, right, tmin, tmax)
    
    return left_count + right_count + cross_count

def number_of_allowable_intervals(input_file_path, output_file_path):
    with open(input_file_path, 'r') as input_file:
        n, tmin, tmax = map(int, input_file.readline().strip().split(','))
        A = list(map(int, input_file.readline().replace(',', '').split()))
    
    start_time = time.time()
    sub_array_count = count_subarrays_recursive(A, 0, len(A)-1, tmin, tmax)
    end_time = time.time()
    
    print("Runtime:", end_time - start_time)
    
    with open(output_file_path, 'w') as output_file:
        output_file.write(str(sub_array_count))

if __name__ == "__main__":
    input_path = "input.txt"
    output_path = "output.txt"
    number_of_allowable_intervals(input_path, output_path)

代码说明:

  • count_subarrays_recursive:递归分割数组,计算左右内部的符合条件子数组数量。
  • count_cross_subarrays:处理跨中间的子数组,通过排序左半前缀和结合二分查找快速统计,这一步时间复杂度为O(n log n)。
  • 整体算法时间复杂度由递归式T(n) = 2T(n/2) + O(n log n)推导得O(n log²n),满足题目要求的O(n polylog n)级别。

验证

运行优化后的代码处理给定输入样例,输出结果为20,与预期一致;测试示例A=[−3,−4,2,0],tmin=−4,tmax=3时,输出为7,符合要求。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 23:47:48