Skip to content

题目链接 · 灵神原题解(署名来源)

前言

本质上,我们要在两个有序数组中,找到第 k 小的数,其中 k=m+n2

  • 如果 m+n 是奇数,返回第 k 小的数。例如 m+n=5,返回第 52=3 小的数。
  • 如果 m+n 是偶数,返回第 k 小的数和第 k+1 小的数的平均值。例如 m+n=6,返回第 3 小的数和第 4 小的数的平均值。

本文先从最暴力的排序做法开始,然后讲解双指针做法,最后过渡到二分做法。

一、引入:均匀分组

lc4-1-c.png

这里的关键是「均匀分组」,每组 5 个数,只要第一组的最大值 第二组的最小值,我们就找到了答案。

怎么想到要均匀分组的?请看百科中关于中位数的介绍:

中位数……可将数值集合划分大小相等的两部分。

mergeda+b 排序后的数组。

k 小的数在 merged 中的下标为 k1,也就是 m+n21=m+n+121=m+n12

python
class Solution:
    def findMedianSortedArrays(self, a: List[int], b: List[int]) -> float:
        merged = a + b
        merged.sort()

        s = len(merged)
        k = (s - 1) // 2
        return merged[k] if s % 2 else (merged[k] + merged[k + 1]) / 2
cpp
// C++ 版待补充
cpp
class Solution {
public:
    double findMedianSortedArrays(vector<int>& a, vector<int>& b) {
        auto merged = a;
        merged.insert(merged.end(), b.begin(), b.end());
        ranges::sort(merged);

        int s = merged.size();
        int k = (s - 1) / 2;
        return s % 2 ? merged[k] : (merged[k] + merged[k + 1]) / 2.0;
    }
};

复杂度分析

  • 时间复杂度:O((m+n)log(m+n)),其中 ma 的长度,nb 的长度。
  • 空间复杂度:O(m+n)

二、枚举:双指针做法

如果您在阅读的过程中产生了一些疑问,请看后文的「答疑」。

lc4-4-c2.png

下面来说具体做法。

ab 的长度分别为 mn,且 mn(如果 m>n 则交换 ab)。

  • 为方便处理 i=1,即 a0 个数在第一组的情况,我们可以往 a 的最左边插入一个哨兵 ,这可以保证数组仍然是有序的。对于 j=1 的情况也同理,往 b 的最左边插入一个
  • 为方便处理 i+1=m,即 am 个数在第一组的情况,我们可以往 a 的最右边插入一个哨兵 ,这可以保证数组仍然是有序的。对于 j+1=n 的情况也同理,往 b 的最右边插入一个 。这可以避免 ai+1bj+1 下标越界。
  • 插入 后,便可保证无论 ab 是什么样的,一定存在一个 i,满足 aibj+1ai+1>bj
  • mn 的值不变。

如此修改后,i 的含义变成了 ai 个数在第一组,j 的含义变成了 bj 个数在第一组。

初始化 i=0,那么 j 应该初始化成多少?

  • 如果 m+n 是偶数,那么每组的大小为 m+n2j 应当初始化成 m+n2
  • 如果 m+n 是奇数,我们规定第一组比第二组多一个数,第一组的大小为 m+n+12j 应当初始化成 m+n+12

两种情况可以合并为:j 初始化成 m+n+12

为了保证组的大小不变,i 每增加 1j 就要减少 1。所以有

j=m+n+12i

根据图片中的结论,只要发现 aibj+1ai+1>bj,那么:

  • 如果 m+n 是偶数,中位数为 max(ai,bj)min(ai+1,bj+1) 的平均值。
  • 如果 m+n 是奇数,中位数为 max(ai,bj)

答疑

:为什么图中说存在一个位置,满足 aibj+1ai+1>bj

:根据 ij 的关系,i 变大,j 会随着变小。把 b 反转,变成一个递减数组,这样 j 会随着 i 的变大而变大,我们可以更容易地观察出性质。把这两个数组画成折线图,一个递增另一个递减,并且由于我们插入了 ,所以二者必然相交。这说明存在一个位置,满足 aibj+1ai+1>bj

:保证 mn 有什么好处?

:如果 m>n,我们没法从 i=0 开始枚举。以 m=5,n=3 为例,i=0 时,b 数组需要有 4 个数在第一组,但 n=3<4,无法做到。保证 mn 可以让我们从 i=0 开始枚举,写起来更方便。

:如果数组中存在重复元素,上述做法是否正确?

:仍然是正确的,因为只用到了「ab 是有序数组」的条件。

写法一

python
class Solution:
    def findMedianSortedArrays(self, a: List[int], b: List[int]) -> float:
        if len(a) > len(b):
            a, b = b, a  # 保证下面的 i 可以从 0 开始枚举

        m, n = len(a), len(b)
        a = [-inf] + a + [inf]
        b = [-inf] + b + [inf]

        # 枚举 nums1 有 i 个数在第一组
        # 那么 nums2 有 j = (m + n + 1) // 2 - i 个数在第一组
        i, j = 0, (m + n + 1) // 2
        while True:
            if a[i] <= b[j + 1] and a[i + 1] > b[j]:  # 写 >= 也可以
                max1 = max(a[i], b[j])  # 第一组的最大值
                min2 = min(a[i + 1], b[j + 1])  # 第二组的最小值
                return max1 if (m + n) % 2 else (max1 + min2) / 2
            i += 1  # 继续枚举
            j -= 1
cpp
// C++ 版待补充
cpp
class Solution {
public:
    double findMedianSortedArrays(vector<int>& a, vector<int>& b) {
        if (a.size() > b.size()) {
            swap(a, b); // 保证下面的 i 可以从 0 开始枚举
        }

        int m = a.size(), n = b.size();
        a.insert(a.begin(), INT_MIN); // 最左边插入 -∞
        b.insert(b.begin(), INT_MIN);
        a.push_back(INT_MAX); // 最右边插入 ∞
        b.push_back(INT_MAX);

        // 枚举 nums1 有 i 个数在第一组
        // 那么 nums2 有 j = (m + n + 1) / 2 - i 个数在第一组
        int i = 0, j = (m + n + 1) / 2;
        while (true) {
            if (a[i] <= b[j + 1] && a[i + 1] > b[j]) { // 写 >= 也可以
                int max1 = max(a[i], b[j]); // 第一组的最大值
                int min2 = min(a[i + 1], b[j + 1]); // 第二组的最小值
                return (m + n) % 2 ? max1 : (max1 + min2) / 2.0;
            }
            i++; // 继续枚举
            j--;
        }
    }
};

写法二(小优化)

实际上,aibj+1ai+1>bj 这两个条件不需要都判断,我们只需要判断其中一个即可。为什么?且听我说。

由于 ab 是有序的,随着 i 的不断变大,j 的不断变小,ai+1>bj 会从「不成立」变成「成立」。

把循环条件改成:如果 ai+1bj,继续循环;如果 ai+1>bj,退出循环。

退出循环之后,除了可以说明 ai+1>bj 成立外,还有一个隐含的性质:在退出循环之前的最后一轮循环,我们在比较哪两个数?正好就是 aibj+1!并且这两个数的大小关系是 aibj+1

python
class Solution:
    def findMedianSortedArrays(self, a: List[int], b: List[int]) -> float:
        if len(a) > len(b):
            a, b = b, a  # 保证下面的 i 可以从 0 开始枚举

        m, n = len(a), len(b)
        a = [-inf] + a + [inf]
        b = [-inf] + b + [inf]

        # 枚举 nums1 有 i 个数在第一组
        # 那么 nums2 有 j = (m + n + 1) // 2 - i 个数在第一组
        i, j = 0, (m + n + 1) // 2
        while a[i + 1] <= b[j]:
            i += 1  # 继续枚举
            j -= 1

        max1 = max(a[i], b[j])  # 第一组的最大值
        min2 = min(a[i + 1], b[j + 1])  # 第二组的最小值
        return max1 if (m + n) % 2 else (max1 + min2) / 2
cpp
// C++ 版待补充
cpp
class Solution {
public:
    double findMedianSortedArrays(vector<int>& a, vector<int>& b) {
        if (a.size() > b.size()) {
            swap(a, b); // 保证下面的 i 可以从 0 开始枚举
        }

        int m = a.size(), n = b.size();
        a.insert(a.begin(), INT_MIN); // 最左边插入 -∞
        b.insert(b.begin(), INT_MIN);
        a.push_back(INT_MAX); // 最右边插入 ∞
        b.push_back(INT_MAX);

        // 枚举 nums1 有 i 个数在第一组
        // 那么 nums2 有 j = (m + n + 1) / 2 - i 个数在第一组
        int i = 0, j = (m + n + 1) / 2;
        while (a[i + 1] <= b[j]) {
            i++; // 继续枚举
            j--;
        }

        int max1 = max(a[i], b[j]); // 第一组的最大值
        int min2 = min(a[i + 1], b[j + 1]); // 第二组的最小值
        return (m + n) % 2 ? max1 : (max1 + min2) / 2.0;
    }
};

复杂度分析

  • 时间复杂度:O(m+n),其中 ma 的长度,nb 的长度。往 a 前面插入一个元素的时间复杂度是 O(m),往 b 前面插入一个元素的时间复杂度是 O(n),加起来是 O(m+n)
  • 空间复杂度:O(m+n)

三、优化:二分做法

由于 ab 是有序数组,i 越小,aibj+1 越能成立;i 越大,aibj+1 越不能成立。

所以可以二分最大的满足 aibj+1i。二分结束后,我们有 aibj+1ai+1>bj

关于二分的原理,请看视频【基础算法精讲 04】

最后,讨论二分的上下界。本文用开区间二分,其他二分写法也是可以的。

  • 开区间二分左边界:0。在插入 后,aibj+1i=0 时一定成立。
  • 开区间二分右边界:m+1。在插入 后,aibj+1i=m+1 时一定不成立。

答疑

:能否二分红色折线图的最小值?

:这种做法会在有重复元素时失效。试想一下,如果我们在折线图上二分,碰巧遇到了相邻且相同的元素,你要更新 left 还是更新 right 呢?

:为什么上面的双指针写法二,循环判断的是 ai+1bj,这里却变成了 aibj+1

:本质是一样的。上面的双指针写法,也可以写成 aibj+1,但需要在循环结束后把 i 减一,把 j 加一。对比来看,上面的双指针写法,循环结束后立刻得到了我们想要的 i,而写成 aibj+1 需要额外做个加减一的微调。下面的二分写法,根据循环不变量的定义,写成 aibj+1,循环结束后的 left 就是我们想要的 i

写法一

注意在数组前面插入元素的时间复杂度是线性的,所以和上面的复杂度分析一样,都是 O(n+m)

真正满足题目时间复杂度要求的是后面的写法二。

python
class Solution:
    def findMedianSortedArrays(self, a: List[int], b: List[int]) -> float:
        if len(a) > len(b):
            a, b = b, a

        m, n = len(a), len(b)
        a = [-inf] + a + [inf]
        b = [-inf] + b + [inf]

        # 循环不变量:a[left] <= b[j+1]
        # 循环不变量:a[right] > b[j+1]
        left, right = 0, m + 1
        while left + 1 < right:  # 开区间 (left, right) 不为空
            i = (left + right) // 2
            j = (m + n + 1) // 2 - i
            if a[i] <= b[j + 1]:
                left = i  # 缩小二分区间为 (i, right)
            else:
                right = i  # 缩小二分区间为 (left, i)

        # 此时 left 等于 right-1
        # a[left] <= b[j+1] 且 a[right] > b[(j-1)+1] = b[j],所以答案是 i=left
        i = left
        j = (m + n + 1) // 2 - i
        max1 = max(a[i], b[j])
        min2 = min(a[i + 1], b[j + 1])
        return max1 if (m + n) % 2 else (max1 + min2) / 2
cpp
// C++ 版待补充
cpp
class Solution {
public:
    double findMedianSortedArrays(vector<int>& a, vector<int>& b) {
        if (a.size() > b.size()) {
            swap(a, b); // 保证下面的 i 可以从 0 开始枚举
        }

        int m = a.size(), n = b.size();
        a.insert(a.begin(), INT_MIN);
        b.insert(b.begin(), INT_MIN);
        a.push_back(INT_MAX);
        b.push_back(INT_MAX);

        // 循环不变量:a[left] <= b[j+1]
        // 循环不变量:a[right] > b[j+1]
        int left = 0, right = m + 1;
        while (left + 1 < right) { // 开区间 (left, right) 不为空
            int i = left + (right - left) / 2;
            int j = (m + n + 1) / 2 - i;
            if (a[i] <= b[j + 1]) {
                left = i; // 缩小二分区间为 (i, right)
            } else {
                right = i; // 缩小二分区间为 (left, i)
            }
        }

        // 此时 left 等于 right-1
        // a[left] <= b[j+1] 且 a[right] > b[(j-1)+1] = b[j],所以答案是 i=left
        int i = left;
        int j = (m + n + 1) / 2 - i;
        int max1 = max(a[i], b[j]);
        int min2 = min(a[i + 1], b[j + 1]);
        return (m + n) % 2 ? max1 : (max1 + min2) / 2.0;
    }
};

写法二(最终版本)

去掉插入的 ab 的下标都减一。

如此修改后,i 的含义变成了 ai+1 个数在第一组,j 的含义变成了 bj+1 个数在第一组。

前文的关系式

j=m+n+12i

修改成

j+1=m+n+12(i+1)

j=m+n+12i2=m+n32i

开区间二分的左右边界改成 1m

答疑

:当 m=0 时,是否会算出 i=0

:不会,m=0 不会进入二分循环,i=left=1

python
class Solution:
    def findMedianSortedArrays(self, a: List[int], b: List[int]) -> float:
        if len(a) > len(b):
            a, b = b, a

        m, n = len(a), len(b)
        # 循环不变量:a[left] <= b[j+1]
        # 循环不变量:a[right] > b[j+1]
        left, right = -1, m
        while left + 1 < right:  # 开区间 (left, right) 不为空
            i = (left + right) // 2
            j = (m + n - 3) // 2 - i
            if a[i] <= b[j + 1]:
                left = i  # 缩小二分区间为 (i, right)
            else:
                right = i  # 缩小二分区间为 (left, i)

        # 此时 left 等于 right-1
        # a[left] <= b[j+1] 且 a[right] > b[(j-1)+1] = b[j],所以答案是 i=left
        i = left
        j = (m + n - 3) // 2 - i
        ai = a[i] if i >= 0 else -inf
        bj = b[j] if j >= 0 else -inf
        ai1 = a[i + 1] if i + 1 < m else inf
        bj1 = b[j + 1] if j + 1 < n else inf
        max1 = max(ai, bj)
        min2 = min(ai1, bj1)
        return max1 if (m + n) % 2 else (max1 + min2) / 2
cpp
// C++ 版待补充
python
class Solution:
    def findMedianSortedArrays(self, a: List[int], b: List[int]) -> float:
        if len(a) > len(b):
            a, b = b, a

        m, n = len(a), len(b)
        # 注意 range(m) 是 O(1) 的,不是 O(m)
        i = bisect_left(range(m), True, key=lambda i: a[i] > b[(m + n - 1) // 2 - i]) - 1

        j = (m + n - 3) // 2 - i
        ai = a[i] if i >= 0 else -inf
        bj = b[j] if j >= 0 else -inf
        ai1 = a[i + 1] if i + 1 < m else inf
        bj1 = b[j + 1] if j + 1 < n else inf
        max1 = max(ai, bj)
        min2 = min(ai1, bj1)
        return max1 if (m + n) % 2 else (max1 + min2) / 2
cpp
// C++ 版待补充
cpp
class Solution {
public:
    double findMedianSortedArrays(vector<int>& a, vector<int>& b) {
        if (a.size() > b.size()) {
            swap(a, b);
        }

        int m = a.size(), n = b.size();
        // 循环不变量:a[left] <= b[j+1]
        // 循环不变量:a[right] > b[j+1]
        int left = -1, right = m;
        while (left + 1 < right) { // 开区间 (left, right) 不为空
            int i = left + (right - left) / 2;
            int j = (m + n + 1) / 2 - i - 2;
            if (a[i] <= b[j + 1]) {
                left = i; // 缩小二分区间为 (i, right)
            } else {
                right = i; // 缩小二分区间为 (left, i)
            }
        }

        // 此时 left 等于 right-1
        // a[left] <= b[j+1] 且 a[right] > b[(j-1)+1] = b[j],所以答案是 i=left
        int i = left;
        int j = (m + n + 1) / 2 - i - 2;
        int ai = i >= 0 ? a[i] : INT_MIN;
        int bj = j >= 0 ? b[j] : INT_MIN;
        int ai1 = i + 1 < m ? a[i + 1] : INT_MAX;
        int bj1 = j + 1 < n ? b[j + 1] : INT_MAX;
        int max1 = max(ai, bj);
        int min2 = min(ai1, bj1);
        return (m + n) % 2 ? max1 : (max1 + min2) / 2.0;
    }
};

复杂度分析

  • 时间复杂度:O(logmin(m,n)),其中 ma 的长度,nb 的长度。:这个复杂度比题目所要求的 O(log(m+n)) 更优!
  • 空间复杂度:O(1)

思考题

改成求两个有序数组的第 k 小,代码要怎么改?

欢迎在评论区分享你的思路/代码。

分类题单

如何科学刷题?

  1. 滑动窗口与双指针(定长/不定长/单序列/双序列/三指针/分组循环)
  2. 二分算法(二分答案/最小化最大值/最大化最小值/第K小)
  3. 单调栈(基础/矩形面积/贡献法/最小字典序)
  4. 网格图(DFS/BFS/综合应用)
  5. 位运算(基础/性质/拆位/试填/恒等式/思维)
  6. 图论算法(DFS/BFS/拓扑排序/基环树/最短路/最小生成树/网络流)
  7. 动态规划(入门/背包/划分/状态机/区间/状压/数位/数据结构优化/树形/博弈/概率期望)
  8. 常用数据结构(前缀和/差分/栈/队列/堆/字典树/并查集/树状数组/线段树)
  9. 数学算法(数论/组合/概率期望/博弈/计算几何/随机算法)
  10. 贪心与思维(基本贪心策略/反悔/区间/字典序/数学/思维/脑筋急转弯/构造)
  11. 链表、二叉树与回溯(前后指针/快慢指针/DFS/BFS/直径/LCA/一般树)
  12. 字符串(KMP/Z函数/Manacher/字符串哈希/AC自动机/后缀数组/子序列自动机)

我的题解精选(已分类)

欢迎关注 B站@灵茶山艾府

本文整理自灵茶山艾府(endlesscheng)的公开内容,仅供个人学习使用

本站仅供个人学习使用,请勿外传