数组所有子数组不平衡值计算算法优化求助
算法优化方案
问题分析
你的原始算法通过枚举所有子数组,排序后计算不平衡值,时间复杂度为O(n³ logn)(枚举子数组是O(n²),每个子数组排序是O(k logk),k最大为n,总复杂度因此拉满),数组长度稍大就会超时。我们需要转换思路,通过计算每个元素对总不平衡值的直接贡献来大幅降低时间复杂度。
核心思路
子数组的不平衡值是排序后相邻元素差值>1的次数总和。换个角度看:对于每个数值x,统计有多少个子数组满足三个条件:
- 子数组包含
x; - 子数组不包含
x+1; - 子数组中存在比
x大的元素。
每个满足条件的子数组,排序后x的下一个元素必然≥x+2,会贡献1次不平衡值。总不平衡值就是所有x对应这类子数组的数量之和。
具体实现步骤
- 记录数组中每个数值的出现位置,方便快速查找
x+1的位置; - 用单调栈预处理每个元素右侧第一个更大元素的位置,快速判断子数组中是否存在比
x大的元素; - 对每个元素
arr[i] = x:- 找到左边最近的
x+1的位置(无则为-1); - 找到右边最近的
x+1的位置(无则为数组长度n); - 计算包含
x且不含x+1的子数组总数:(i - 左边界可选数) * (右边界可选数 - i); - 从中减去包含
x、不含x+1且没有比x大的元素的子数组数量,得到该元素对总不平衡值的贡献;
- 找到左边最近的
- 累加所有元素的贡献,得到最终结果。
优化后代码示例
def func(arr): n = len(arr) if n <= 1: return 0 # 记录每个数值的所有出现位置 num_pos = dict() for idx, num in enumerate(arr): if num not in num_pos: num_pos[num] = [] num_pos[num].append(idx) # 预处理每个位置右侧第一个更大元素的索引(单调栈实现) next_greater = [n] * n stack = [] for i in range(n-1, -1, -1): while stack and arr[stack[-1]] <= arr[i]: stack.pop() if stack: next_greater[i] = stack[-1] stack.append(i) total_imbalance = 0 for i in range(n): x = arr[i] x_plus_1 = x + 1 # 找左侧最近的x+1位置 left_bound = -1 if x_plus_1 in num_pos: positions = num_pos[x_plus_1] low, high = 0, len(positions)-1 best = -1 while low <= high: mid = (low + high) // 2 if positions[mid] < i: best = positions[mid] low = mid + 1 else: high = mid - 1 left_bound = best # 找右侧最近的x+1位置 right_bound = n if x_plus_1 in num_pos: positions = num_pos[x_plus_1] low, high = 0, len(positions)-1 best = n while low <= high: mid = (low + high) // 2 if positions[mid] > i: best = positions[mid] high = mid - 1 else: low = mid + 1 right_bound = best # 计算当前元素的贡献值 total_valid = (i - left_bound) * (right_bound - i) no_greater = (i - left_bound) * (min(next_greater[i], right_bound) - i) total_imbalance += (total_valid - no_greater) return total_imbalance
复杂度分析
- 单调栈预处理下一个更大元素:O(n);
- 每个元素查找左右
x+1位置:O(logm)(m为x+1的出现次数),总复杂度O(n logn); - 整体时间复杂度为O(n logn),相比原始算法的O(n³ logn),效率提升几个数量级,可处理大规模数组。
内容的提问来源于stack exchange,提问作者Sharhad
相关产品推荐
相关产品推荐

