寻找两个正序数组的中位数 Java 题解,二分分割详解

4:寻找两个正序数组的中位数 Java 题解,二分分割详解

给定两个正序数组,要求在 O(log(m+n)) 时间内求出整体中位数。本文从直观的合并方法入手,重点讲解如何在较短数组上二分查找分割线,并附可直接提交的 Java 答案、完整测试代码和测试结果。

一、题目描述

给定两个大小分别为 mn 的正序数组 nums1nums2,找出并返回这两个正序数组的中位数。

要求算法的时间复杂度为 O(log(m+n))

示例 1:

text 复制代码
输入:nums1 = [1,3], nums2 = [2]
输出:2.00000
解释:合并后为 [1,2,3],中位数是 2。

示例 2:

text 复制代码
输入:nums1 = [1,2], nums2 = [3,4]
输出:2.50000
解释:合并后为 [1,2,3,4],中位数是 (2 + 3) / 2 = 2.5。

二、中位数是什么

将所有数字按从小到大排列:

  • 总长度为奇数时,中位数是正中间的数字;
  • 总长度为偶数时,中位数是中间两个数字的平均值。

如果先合并两个数组,再直接取中间位置,时间复杂度是 O(m+n)。这种方法容易理解,但没有满足题目要求的对数级复杂度。

三、核心思路:把两个数组同时分成左右两半

分别在两个数组中放置一条分割线:

text 复制代码
nums1: [ 左半部分 | 右半部分 ]
nums2: [ 左半部分 | 右半部分 ]

希望分割满足两个条件:

  1. 两个数组左半部分的元素总数,等于或只比右半部分多一个;
  2. 左半部分的所有元素都不大于右半部分的所有元素。

第二个条件只需要比较分割线附近的四个值:

text 复制代码
maxLeft1  <= minRight2
maxLeft2  <= minRight1

当这两个不等式同时成立时,就找到了正确分割。

四、为什么只在较短数组上二分

假设 nums1 的长度为 mnums2 的长度为 n,并保证 m <= n

nums1 中选择分割位置 partition1,那么 nums2 的分割位置可以直接计算:

java 复制代码
int partition2 = (m + n + 1) / 2 - partition1;

因为左半区总元素数量已经确定,只需要搜索一个分割位置。

在较短数组上二分有两个好处:

  • 时间复杂度为 O(log(min(m, n)))
  • 可以保证计算出的 partition2 不会越过较长数组的有效范围。

五、推荐答案:Java 二分实现

java 复制代码
class Solution {
    public double findMedianSortedArrays(int[] nums1, int[] nums2) {
        if (nums1.length > nums2.length) {
            return findMedianSortedArrays(nums2, nums1);
        }

        int m = nums1.length;
        int n = nums2.length;
        int low = 0;
        int high = m;

        while (low <= high) {
            int partition1 = low + (high - low) / 2;
            int partition2 = (m + n + 1) / 2 - partition1;

            int maxLeft1 = partition1 == 0
                    ? Integer.MIN_VALUE : nums1[partition1 - 1];
            int minRight1 = partition1 == m
                    ? Integer.MAX_VALUE : nums1[partition1];

            int maxLeft2 = partition2 == 0
                    ? Integer.MIN_VALUE : nums2[partition2 - 1];
            int minRight2 = partition2 == n
                    ? Integer.MAX_VALUE : nums2[partition2];

            if (maxLeft1 <= minRight2 && maxLeft2 <= minRight1) {
                if ((m + n) % 2 == 1) {
                    return Math.max(maxLeft1, maxLeft2);
                }

                long leftMax = Math.max(maxLeft1, maxLeft2);
                long rightMin = Math.min(minRight1, minRight2);
                return (leftMax + rightMin) / 2.0;
            }

            if (maxLeft1 > minRight2) {
                high = partition1 - 1;
            } else {
                low = partition1 + 1;
            }
        }

        throw new IllegalArgumentException("输入数组不是正序数组");
    }
}

六、四个边界值如何理解

每条分割线两侧各有一个关键值:

text 复制代码
nums1: ... maxLeft1 | minRight1 ...
nums2: ... maxLeft2 | minRight2 ...

只要满足:

java 复制代码
maxLeft1 <= minRight2 && maxLeft2 <= minRight1

就可以确定整个左半区不大于整个右半区。

如果分割线位于数组最左侧,左侧没有元素,可以将左侧最大值视为负无穷;如果位于最右侧,右侧没有元素,可以将右侧最小值视为正无穷:

java 复制代码
int maxLeft1 = partition1 == 0
        ? Integer.MIN_VALUE : nums1[partition1 - 1];
int minRight1 = partition1 == m
        ? Integer.MAX_VALUE : nums1[partition1];

这样可以统一处理空数组和边界分割,不需要增加多层特殊判断。

七、如何决定二分方向

如果:

java 复制代码
maxLeft1 > minRight2

说明 nums1 左侧取了太多较大的元素,分割线需要左移:

java 复制代码
high = partition1 - 1;

否则说明 nums1 左侧元素太少,分割线需要右移:

java 复制代码
low = partition1 + 1;

八、奇数和偶数长度如何取中位数

总长度为奇数

通过 (m + n + 1) / 2 的设计,左半区会比右半区多一个元素。因此中位数就是左半区的最大值:

java 复制代码
return Math.max(maxLeft1, maxLeft2);

总长度为偶数

中位数是左半区最大值与右半区最小值的平均值:

java 复制代码
long leftMax = Math.max(maxLeft1, maxLeft2);
long rightMin = Math.min(minRight1, minRight2);
return (leftMax + rightMin) / 2.0;

这里先转换为 long 再相加,是为了避免两个较大的 int 相加时发生整数溢出。

九、逐步分析示例 2

输入:

text 复制代码
nums1 = [1, 2]
nums2 = [3, 4]

partition1 = 2partition2 = 0 时:

text 复制代码
maxLeft1  = 2
minRight1 = +∞
maxLeft2  = -∞
minRight2 = 3

满足:

text 复制代码
2 <= 3
-∞ <= +∞

总长度为偶数,因此:

text 复制代码
(max(2, -∞) + min(+∞, 3)) / 2
= (2 + 3) / 2
= 2.5

十、复杂度分析

假设较短数组的长度为 min(m, n)

  • 时间复杂度:O(log(min(m, n)))
  • 空间复杂度:O(1)

这比合并数组的 O(m+n) 更符合题目要求。

十一、完整可运行代码与测试用例

java 复制代码
public class MedianSortedArraysTest {

    public static double findMedianSortedArrays(int[] nums1, int[] nums2) {
        if (nums1.length > nums2.length) {
            return findMedianSortedArrays(nums2, nums1);
        }

        int m = nums1.length;
        int n = nums2.length;
        int low = 0;
        int high = m;

        while (low <= high) {
            int partition1 = low + (high - low) / 2;
            int partition2 = (m + n + 1) / 2 - partition1;

            int maxLeft1 = partition1 == 0
                    ? Integer.MIN_VALUE : nums1[partition1 - 1];
            int minRight1 = partition1 == m
                    ? Integer.MAX_VALUE : nums1[partition1];
            int maxLeft2 = partition2 == 0
                    ? Integer.MIN_VALUE : nums2[partition2 - 1];
            int minRight2 = partition2 == n
                    ? Integer.MAX_VALUE : nums2[partition2];

            if (maxLeft1 <= minRight2 && maxLeft2 <= minRight1) {
                if ((m + n) % 2 == 1) {
                    return Math.max(maxLeft1, maxLeft2);
                }

                long leftMax = Math.max(maxLeft1, maxLeft2);
                long rightMin = Math.min(minRight1, minRight2);
                return (leftMax + rightMin) / 2.0;
            }

            if (maxLeft1 > minRight2) {
                high = partition1 - 1;
            } else {
                low = partition1 + 1;
            }
        }

        throw new IllegalArgumentException("输入数组不是正序数组");
    }

    private static void runTest(
            int number,
            int[] nums1,
            int[] nums2,
            double expected) {

        double actual = findMedianSortedArrays(nums1, nums2);
        boolean passed = Math.abs(actual - expected) < 1e-9;

        System.out.println("测试用例 " + number);
        System.out.println("expected = " + expected);
        System.out.println("actual   = " + actual);
        System.out.println("result   = " + (passed ? "PASS" : "FAIL"));
        System.out.println();
    }

    public static void main(String[] args) {
        runTest(1, new int[]{1, 3}, new int[]{2}, 2.0);
        runTest(2, new int[]{1, 2}, new int[]{3, 4}, 2.5);
        runTest(3, new int[]{}, new int[]{1}, 1.0);
        runTest(4, new int[]{0, 0}, new int[]{0, 0}, 0.0);
        runTest(5,
                new int[]{Integer.MAX_VALUE},
                new int[]{Integer.MAX_VALUE},
                Integer.MAX_VALUE);
    }
}

第五组用例专门验证偶数长度求平均值时不会发生整数溢出。

十二、测试结果

实际运行结果如下:

text 复制代码
测试用例 1
expected = 2.0
actual   = 2.0
result   = PASS

测试用例 2
expected = 2.5
actual   = 2.5
result   = PASS

测试用例 3
expected = 1.0
actual   = 1.0
result   = PASS

测试用例 4
expected = 0.0
actual   = 0.0
result   = PASS

测试用例 5
expected = 2.147483647E9
actual   = 2.147483647E9
result   = PASS

五组测试全部通过。

十三、常见错误

1. 直接合并两个数组

合并法能够得到正确结果,但时间复杂度是 O(m+n),不满足题目要求。

2. 没有保证第一个数组更短

如果直接在较长数组上二分,另一个数组的分割位置可能越界。应先交换数组,保证 nums1.length <= nums2.length

3. 左半区数量计算错误

使用:

java 复制代码
(m + n + 1) / 2

其中的 +1 可以让奇数总长度时多出的元素自然落在左半区。

4. 分割线边界发生越界

分割线位于数组两端时,应使用 Integer.MIN_VALUEInteger.MAX_VALUE 作为哨兵值。

5. 平均值发生整数溢出

不要直接计算:

java 复制代码
(leftMax + rightMin) / 2.0

如果两个变量都是 int,加法会先按照整数计算,可能在转换成 double 前就已经溢出。应先使用 long 保存两个边界值。

十四、总结

这道题的难点是从"合并数组"切换到"寻找正确分割线":

  1. 始终在较短数组上进行二分;
  2. 根据左半区总长度计算另一个数组的分割位置;
  3. 使用四个分割边界判断左右区域是否有序;
  4. 根据总长度的奇偶性计算中位数;
  5. 使用 long 避免偶数长度取平均值时溢出。

理解分割线后,这道困难题会从复杂的数组合并问题,变成一个标准的二分查找问题。

更多技术实践与案例,可在 📗【姚前述】中查阅。

相关推荐
半亩码田1 小时前
C#转Python第3.1篇:Python 的 class 没有访问修饰符?面向对象的另一条路
开发语言·python·c#
真上帝的左手2 小时前
10. 软件设计&架构-Spring Security 7 整合CAS SSO 单点登录
java·spring·架构·sso
Wzx1980122 小时前
python沙箱和docker沙箱你选对了吗?
开发语言·python·docker
牧羊人.3332 小时前
Python 办公自动化从入门到入土|09 数据容器之字典
开发语言·python
Zane19942 小时前
单例线程安全、生产者消费者、死锁:并发面试三连问串讲
java·后端
小僧景贤2 小时前
嵌入式C语言 第二篇:基础语法|嵌入式C与标准C的核心差异
c语言·开发语言·嵌入式c语言
阿pin3 小时前
Android随笔-AIDL
android·开发语言·aidl
云烟成雨TD3 小时前
Micrometer 系列【42】链路追踪:Span 体系 | 核心 API
java·链路追踪·micrometer
Gorway3 小时前
理解 Spring 依赖注入:从构造器注入到集合与条件 Bean
java·后端