LeetCode 4. 寻找两个正序数组的中位数
本节目标
将中位数转化为第 k 小元素,用有序前缀排除在两个数组间缩小排名。
这道题是分治与二分综合框架中的拓展题。两个数组分别有序,但合并后的中间位置同时依赖两边;我们不真正合并,而是直接寻找合并序列的第 k 小元素。
题意与约束
给定两个非递减整数数组,返回它们合并后的中位数。总长度为奇数时,中位数是第 (n + 1) / 2 小;为偶数时,是第 n / 2 小和第 n / 2 + 1 小的平均值。数组可能为空,但不会同时为空。
朴素思路与瓶颈
完整合并能在 O(m + n) 时间和 O(m + n) 额外空间内得到答案;双指针只走到中间虽可省空间,最坏仍是线性时间。两者都没有利用“只需要固定排名”,因此本题改为每轮排除一段有序前缀。
将中位数转化为第 k 小
设两个未排除部分从 start1、start2 开始,当前目标为两部分中的第 k 小。若某一数组已耗尽,答案直接是另一数组的第 k 个剩余元素;若 k = 1,答案就是两边当前首元素的较小者。
其余情况令 half = k / 2,比较两边各自第 half 个剩余候选。数组不足 half 个元素时,将该候选看作正无穷,表示这边没有足够长的前缀可安全排除。
为什么较小候选的前缀可以排除
把两个未排除部分看成一个保留重复副本的多重集,记当前第 k 小值为 v。假设 nums1 的第 half 个候选 x 不大于 nums2 的对应候选。若 v < x,两个数组各自至多有 half - 1 个元素不大于 v,合计少于 k,与 v 的排名矛盾;所以 x <= v,nums1 的这段前缀全部不大于目标值。
删除这 half 个不大于目标值的副本后,排名同步减少为 k - half。若 x < v,删掉的副本都严格排在 v 前面;若 x = v,剩余多重集中严格小于 v 的元素至多来自另一个数组候选前的 half - 1 个,而 k - half >= half,同时剩余的不大于 v 的副本数仍至少为 k - half,因此该排名的值仍是 v。候选相等时对称地删除任一侧的 half 个副本都安全。每轮至少丢弃一个元素,k 严格减小。
代码实现
两份源码都用递归辅助函数保存当前起点与排名。C++ 对偶数中位数先把两个第 k 小结果转为 long long,再除以 2.0,避免整数相加先溢出。
- C++
- Python
#include <algorithm>
#include <limits>
#include <vector>
using namespace std;
class Solution {
private:
int kth(
const vector<int>& nums1,
int start1,
const vector<int>& nums2,
int start2,
int k
) {
if (start1 == static_cast<int>(nums1.size())) {
return nums2[start2 + k - 1];
}
if (start2 == static_cast<int>(nums2.size())) {
return nums1[start1 + k - 1];
}
if (k == 1) {
return min(nums1[start1], nums2[start2]);
}
int half = k / 2;
long long candidate1 = start1 + half - 1 < static_cast<int>(nums1.size())
? nums1[start1 + half - 1]
: numeric_limits<long long>::max();
long long candidate2 = start2 + half - 1 < static_cast<int>(nums2.size())
? nums2[start2 + half - 1]
: numeric_limits<long long>::max();
if (candidate1 <= candidate2) {
return kth(nums1, start1 + half, nums2, start2, k - half);
}
return kth(nums1, start1, nums2, start2 + half, k - half);
}
public:
double findMedianSortedArrays(vector<int>& nums1, vector<int>& nums2) {
int total = static_cast<int>(nums1.size() + nums2.size());
if (total % 2 == 1) {
return kth(nums1, 0, nums2, 0, total / 2 + 1);
}
long long left = kth(nums1, 0, nums2, 0, total / 2);
long long right = kth(nums1, 0, nums2, 0, total / 2 + 1);
return (left + right) / 2.0;
}
};
class Solution:
def _kth(self, nums1, start1, nums2, start2, k):
if start1 == len(nums1):
return nums2[start2 + k - 1]
if start2 == len(nums2):
return nums1[start1 + k - 1]
if k == 1:
return min(nums1[start1], nums2[start2])
half = k // 2
candidate1 = (
nums1[start1 + half - 1]
if start1 + half - 1 < len(nums1)
else float("inf")
)
candidate2 = (
nums2[start2 + half - 1]
if start2 + half - 1 < len(nums2)
else float("inf")
)
if candidate1 <= candidate2:
return self._kth(nums1, start1 + half, nums2, start2, k - half)
return self._kth(nums1, start1, nums2, start2 + half, k - half)
def findMedianSortedArrays(self, nums1, nums2):
total = len(nums1) + len(nums2)
if total % 2 == 1:
return float(self._kth(nums1, 0, nums2, 0, total // 2 + 1))
left = self._kth(nums1, 0, nums2, 0, total // 2)
right = self._kth(nums1, 0, nums2, 0, total // 2 + 1)
return (left + right) / 2.0
复杂度分析
- 时间复杂度:
O(log(m + n))。每轮把目标排名至少缩小约一半。 - 空间复杂度:
O(log(m + n))。递归调用深度与排名缩小次数相同;不创建合并数组。
易错点
- 把不足
k / 2的数组当作已经耗尽;它仍可在后续k = 1时提供答案。 - 偶数总长度只找一个第 k 小,或用整数除法丢掉
.5。 - C++ 先用
int相加两个极值,再转换为浮点数。 - 某一数组耗尽后仍访问它的当前起点。
模式迁移
当多个有序来源共同决定固定排名时,优先寻找“一次能证明排除一个有序前缀”的比较。若来源数量增多或需要持续查询,可以转向堆维护多路候选;本题的关键仍是一次性排名缩减,而非维护完整合并结果。