给定有序数组经三次函数转换后的排序优化方案求解
三次函数转换后排序数组的O(n)复杂度解法
存在比暴力O(nlogn)解法更优的O(n)时间复杂度方案,空间复杂度为O(n)。
核心原理
- 三次函数
f(x) = a*x³ + b*x² + c*x + d的一阶导数f’(x) = 3a*x² + 2b*x + c是二次函数,最多存在2个不同实根,即原三次函数最多有2个驻点,单调性最多切换2次。 - 输入的
nums是预排序数组,代入三次函数得到的y序列最多被2个驻点分割为3个独立单调子序列,每个子序列要么升序要么降序。 - 合并最多3个有序序列的时间复杂度为O(n),远低于全量排序的O(nlogn)。
具体实现步骤
- 计算驻点:求解一阶导数
f’(x)=0的实根,按从小到大排序得到两个驻点(如果无实根或只有一个实根,说明单调段数量不超过2)。 - 分割单调子序列:遍历预排序的
nums,按驻点位置把数组划分为最多3段,每一段内的f(x)单调性一致。 - 预处理为升序序列:判断每个子序列的单调性,若为降序则直接反转得到升序序列。
- 多指针合并有序序列:用多指针法合并最多3个升序子序列,得到最终有序结果。
边界情况处理
- 当
a=0时,三次函数退化为二次/一次/常数函数,最多1个驻点,单调段不超过2个,上述逻辑依然适用,时间复杂度保持O(n)。 - 当驻点落在
nums的取值范围之外时,整个y序列是单一单调序列,直接计算后按需反转即可得到结果,无需合并操作。
代码示例(Python)
def sort_transformed_array(nums, a, b, c, d): def f(x): return a * x ** 3 + b * x ** 2 + c * x + d # 求一阶导数的实根:3a x² + 2b x + c = 0 roots = [] if a != 0: delta = (2*b)**2 - 4 * 3*a * c if delta >= 0: sqrt_delta = delta ** 0.5 r1 = (-2*b - sqrt_delta) / (6 * a) r2 = (-2*b + sqrt_delta) / (6 * a) roots = sorted([r1, r2]) # 按驻点分割原数组为多个单调段 segments = [] prev_idx = 0 for r in roots: split_idx = 0 while split_idx < len(nums) and nums[split_idx] <= r: split_idx += 1 if split_idx > prev_idx: segments.append(nums[prev_idx:split_idx]) prev_idx = split_idx if prev_idx < len(nums): segments.append(nums[prev_idx:]) # 将每个段的计算结果转换为升序序列 sorted_segments = [] for seg in segments: ys = [f(x) for x in seg] if len(ys) >= 2: # 取区间中点的导数符号判断单调性,避免边界误差 mid_x = seg[len(seg)//2] deriv = 3*a*mid_x**2 + 2*b*mid_x + c if deriv < 0: ys = ys[::-1] sorted_segments.append(ys) # 合并最多3个升序序列 res = [] pointers = [0] * len(sorted_segments) while True: min_val = float('inf') min_seg_idx = -1 for i in range(len(sorted_segments)): if pointers[i] < len(sorted_segments[i]) and sorted_segments[i][pointers[i]] < min_val: min_val = sorted_segments[i][pointers[i]] min_seg_idx = i if min_seg_idx == -1: break res.append(min_val) pointers[min_seg_idx] += 1 return res
内容的提问来源于stack exchange,提问作者Neo Staff
相关产品推荐
相关产品推荐

