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

给定有序数组经三次函数转换后的排序优化方案求解

三次函数转换后排序数组的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)。

具体实现步骤

  1. 计算驻点:求解一阶导数f’(x)=0的实根,按从小到大排序得到两个驻点(如果无实根或只有一个实根,说明单调段数量不超过2)。
  2. 分割单调子序列:遍历预排序的nums,按驻点位置把数组划分为最多3段,每一段内的f(x)单调性一致。
  3. 预处理为升序序列:判断每个子序列的单调性,若为降序则直接反转得到升序序列。
  4. 多指针合并有序序列:用多指针法合并最多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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.07 00:57:02