求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)要求。
具体步骤:
- 递归分割:将数组从中间分为左右两半,递归计算左右各自的符合条件子数组数量。
- 统计跨中间的子数组:
- 计算左半部分从中间位置向左的所有前缀和(包含中间点)。
- 计算右半部分从中间位置向右的所有前缀和(包含中间点的下一个位置)。
- 将左半前缀和排序,对右半每个前缀和
s_r,用二分查找统计左半中满足tmin - s_r ≤ s_l ≤ tmax - s_r的s_l数量。
- 合并结果:将左右内部计数与跨中间计数相加,得到总结果。
优化后的代码实现
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
相关产品推荐
相关产品推荐

