4. 寻找两个正序数组的中位数 (Hard)

专题归类: 09-二分查找 · 03-数组与矩阵 LeetCode 链接: https://leetcode.cn/problems/median-of-two-sorted-arrays/


在线做题: 写 Python 代码并运行 · Java 跳转 LeetCode

题目描述

给定两个大小分别为 mn 的正序(从小到大)数组 nums1nums2。请你找出并返回这两个正序数组的中位数。

算法的时间复杂度应该为 O(log(m + n))。

示例 1:

输入:nums1 = [1,3], nums2 = [2]
输出:2.00000
解释:合并数组 = [1,2,3],中位数 2

示例 2:

输入:nums1 = [1,2], nums2 = [3,4]
输出:2.50000
解释:合并数组 = [1,2,3,4],中位数 (2 + 3) / 2 = 2.5

提示:

  • nums1.length == m
  • nums2.length == n
  • 0 <= m <= 1000
  • 0 <= n <= 1000
  • 1 <= m + n <= 2000
  • -10^6 <= nums1[i], nums2[i] <= 10^6

题目详细分析

数据范围含义:

  • m + n <= 2000:O(m+n) 的归并法也能过,但题目要求 O(log(m+n)),暗示必须用二分。
  • 数组可能为空,需要处理空数组的情况。

中位数的定义:

  • 如果总元素数为奇数,中位数 = 第 (m+n+1)//2 小的元素。
  • 如果总元素数为偶数,中位数 = 第 (m+n)//2 小和第 (m+n)//2+1 小的元素的平均值。

问题本质:

  • 等价于”在两个有序数组中找到第 k 小的元素”,k = (m+n+1)//2。
  • 要求 O(log(m+n)),自然想到用二分法排除元素。
  • 关键思想:每次比较两个数组的第 k//2 个元素,排除较小的那一段。

核心约束——在短数组上二分:

  • 必须确保在较短的数组上二分分割线,否则可能出现数组越界。
  • 时间复杂度 O(log(min(m, n))),确保高效。

小白版直白理解

就像两个班级的成绩单(都是按分数从低到高排好的),你要把两份成绩单合并在一起,找出”站在最中间”的那个人的分数。

一种巧妙的方法是:把两个班想象成两把尺子,要找正中间的那条线(分割线)。这条线把两个班的人分成左右两部分,左边的人数和右边的人数相等(或差 1)。线上面的数是左边最大和右边最小,中位数就是由它们算出来的。

比如 [1,2] 和 [3,4]:

  • 在 [1,2] 中切一刀:左边 [1],右边 [2]
  • 在 [3,4] 中切一刀:左边 [3],右边 [4]
  • 左边最大值 = max(1,3) = 3,右边最小值 = min(2,4) = 2
  • 中位数 = (3+2)/2 = 2.5

解题思路

思路一:在短数组上二分分割线(推荐)

核心思想:

  • 在较短的数组上确定一个分割点 i,使得 nums1 的前 i 个元素 + nums2 的前 j 个元素 = 总元素数的一半。
  • 其中 j = (m+n+1)//2 - i。
  • 分割线是否合法:nums1[i-1] <= nums2[j]nums2[j-1] <= nums1[i]
  • 如果不合法,根据比较结果调整 i。

可视化(以 nums1=[1,3,5,7], nums2=[2,4,6,8] 为例):

总元素数 8,一半为 4
在 nums1 中尝试 i=2:
  nums1: [1,3 | 5,7]
  nums2: [2,4 | 6,8]
  左半: [1,3,2,4], 右半: [5,7,6,8]
  检查: nums1[1]=3 <= nums2[2]=6 ✓
        nums2[1]=4 <= nums1[2]=5 ✓
  中位数 = (max(3,4) + min(5,6)) / 2 = (4+5)/2 = 4.5

处理边界:

  • 当 i = 0 时,nums1 左半没有元素,用 -inf 表示左边最大值。
  • 当 i = m 时,nums1 右半没有元素,用 +inf 表示右边最小值。
  • j 同理。
def findMedianSortedArrays(nums1, nums2):
    # 保证 nums1 是较短的数组
    if len(nums1) > len(nums2):
        nums1, nums2 = nums2, nums1
 
    m, n = len(nums1), len(nums2)
    left, right = 0, m
    total_left = (m + n + 1) // 2  # 左半部分需要的元素总数
 
    while left <= right:
        # i 是 nums1 的分割线位置(左侧有 i 个元素)
        i = (left + right) // 2
        # j 是 nums2 的分割线位置(左侧有 j 个元素)
        j = total_left - i
 
        # 处理边界值(用无穷大/小表示空侧)
        nums1_left_max = nums1[i - 1] if i > 0 else float('-inf')
        nums1_right_min = nums1[i] if i < m else float('inf')
        nums2_left_max = nums2[j - 1] if j > 0 else float('-inf')
        nums2_right_min = nums2[j] if j < n else float('inf')
 
        # 检查分割线是否合法
        if nums1_left_max <= nums2_right_min and nums2_left_max <= nums1_right_min:
            # 找到正确分割线
            if (m + n) % 2 == 0:
                # 偶数:取左半最大值和有半最小值的平均值
                return (max(nums1_left_max, nums2_left_max) +
                        min(nums1_right_min, nums2_right_min)) / 2
            else:
                # 奇数:中位数就是左半部分的最大值
                return max(nums1_left_max, nums2_left_max)
        elif nums1_left_max > nums2_right_min:
            # nums1 左半太大,需要缩小 i
            right = i - 1
        else:
            # nums2 左半太大,需要增大 i
            left = i + 1
 
    return 0.0  # 理论上不会执行到这里

思路二:找第 k 小数法

核心思想: 在两个有序数组中找第 k 小的元素。每次比较两个数组的第 k//2 个元素,排除较小的那一半。

def findMedianSortedArrays(nums1, nums2):
    def get_kth(k):
        """返回两个有序数组中第 k 小的元素(k 从 1 开始)"""
        idx1, idx2 = 0, 0
        m, n = len(nums1), len(nums2)
 
        while True:
            # 处理边界:某个数组已空
            if idx1 == m:
                return nums2[idx2 + k - 1]
            if idx2 == n:
                return nums1[idx1 + k - 1]
            if k == 1:
                return min(nums1[idx1], nums2[idx2])
 
            # 比较两个数组的第 k//2 个元素
            half = k // 2
            new_idx1 = min(idx1 + half, m) - 1
            new_idx2 = min(idx2 + half, n) - 1
 
            if nums1[new_idx1] <= nums2[new_idx2]:
                # 排除 nums1 的前 half 个元素
                k -= (new_idx1 - idx1 + 1)
                idx1 = new_idx1 + 1
            else:
                # 排除 nums2 的前 half 个元素
                k -= (new_idx2 - idx2 + 1)
                idx2 = new_idx2 + 1
 
    total = len(nums1) + len(nums2)
    if total % 2 == 1:
        return get_kth(total // 2 + 1)
    else:
        return (get_kth(total // 2) + get_kth(total // 2 + 1)) / 2

时间复杂度: O(log(min(m, n)))(思路一)或 O(log(m+n))(思路二,实际也退化为 O(log(min(m,n))))。 空间复杂度: O(1)。


易错点

  1. 边界值处理:分割线在数组两端时,用 -inf+inf 表示是优雅的处理方式。忘记处理边界会导致索引越界。
  2. 总元素数奇偶性:奇数时中位数是左半最大值,偶数时是左右半的平均值。要把两种情况都考虑到。
  3. 分割线位置 i 的范围i[0, m] 范围内(包括 0 和 m),对应的 j 由 total_left - i 计算。需要确保 0 <= j <= n,这依赖于 nums1 是短数组。
  4. 死循环:在二分 i 的过程中,当 i 需要调整时,必须用 mid +/- 1 更新,不能直接用 mid,否则可能死循环。
  5. total_left 的计算(m+n+1)//2 中的 +1 是为了在总数为奇数时让左半多一个元素,这样中位数就是左半最大值。

框架提炼

两个有序数组找中位数(分割线法)模板:

def findMedianSortedArrays(nums1, nums2):
    # 确保 nums1 是短数组
    if len(nums1) > len(nums2):
        nums1, nums2 = nums2, nums1
 
    m, n = len(nums1), len(nums2)
    left, right = 0, m
    total_left = (m + n + 1) // 2
 
    while left <= right:
        i = left + (right - left) // 2  # nums1 分割线
        j = total_left - i               # nums2 分割线
 
        # 四个边界值
        left1 = nums1[i - 1] if i > 0 else float('-inf')
        right1 = nums1[i] if i < m else float('inf')
        left2 = nums2[j - 1] if j > 0 else float('-inf')
        right2 = nums2[j] if j < n else float('inf')
 
        if left1 <= right2 and left2 <= right1:
            # 找到正确分割线
            if (m + n) % 2 == 0:
                return (max(left1, left2) + min(right1, right2)) / 2
            else:
                return max(left1, left2)
        elif left1 > right2:
            right = i - 1  # i 太大,左移
        else:
            left = i + 1   # i 太小,右移

核心技术要点总结:

  1. 在短数组上二分:保证 j 不越界,且将复杂度降到 O(log(min(m,n))).
  2. 奇偶统一处理total_left = (m+n+1)//2 巧妙处理了奇偶两种情况。
  3. 无穷大/小处理边界:用 ±inf 优雅处理分割线在数组两端的情况。
  4. 四个边界值验证left1 <= right2 and left2 <= right1 是分割线合法的充要条件。

关联题目