一句话说明核心方法
把"求中位数"转化为"求第 k 小的元素",用二分法每次从两个数组中淘汰 k/2 个元素,递归缩小 k,直到 k == 1 时取两个数组头部的较小值。
思路推导
题意转化 :中位数只取决于合并后长度 N = m + n 的奇偶。
- N 为奇数:中位数就是第
(N+1)/2小; - N 为偶数:中位数是第
N/2小与第N/2+1小的平均值。
关键观察 :我们不需要真的把两个数组合并。要在两个有序数组中找第 k 小,只需要比较双方各自第 k/2 个元素:
- 若
nums1[i+half-1] < nums2[j+half-1],说明 nums1 的前 half 个元素一定不可能是第 k 小 ------因为比它们小的元素最多只有(half-1) + (half-1) = k-2个。这 half 个元素可以整段淘汰。 - 淘汰后 k 减去 half,在剩余部分继续找。
算法选择 :每次至少淘汰 k/2 个元素,k 线性衰减到 1,因此是 O(log k) = O(log(m+n)),正好满足题目要求。
统一奇偶的技巧 :令 left = (m+n+1)/2、right = (m+n+2)/2。
- N = 3 时:
left = 2,right = 2,两者相同; - N = 4 时:
left = 2,right = 3,正好是中间两个位置。
用一个公式覆盖两种情况,答案永远写成 (findKth(left) + findKth(right)) / 2.0,不需要写 if/else 分支。
与"归并"的区别:归并是 O(m+n) 时间 + O(m+n) 空间;本题的二分不落地合并数组,空间只有递归栈。
过程示意(findKth 的裁剪过程)
以 nums1 = [1,3]、nums2 = [2] 为例,N = 3,left = right = 2,只需找第 2 小:
findKth(i=0, j=0, k=2)
half = 1
midVal1 = nums1[0] = 1
midVal2 = nums2[0] = 2
1 < 2 → 淘汰 nums1 的前 1 个元素
→ findKth(i=1, j=0, k=1)
k == 1 → min(nums1[1]=3, nums2[0]=2) = 2
以 nums1 = [1,2]、nums2 = [3,4] 为例,N = 4,left = 2、right = 3:
findKth(i=0, j=0, k=2) findKth(i=0, j=0, k=3)
half=1, midVal1=1, midVal2=3 half=1, midVal1=1, midVal2=3
1 < 3 → 淘汰 nums1[0] 1 < 3 → 淘汰 nums1[0]
→ findKth(i=1, j=0, k=1) → findKth(i=1, j=0, k=2)
min(nums1[1]=2, nums2[0]=3) = 2 half=1, midVal1=2, midVal2=3
2 < 3 → 淘汰 nums1[1]
→ findKth(i=2, j=0, k=1)
i 越界 → nums2[0+1-1] = 3
答案 = (2 + 3) / 2.0 = 2.5
可以看到 k 每次都减半,递归深度极浅。
Java 完整代码
java
class Solution {
public double findMedianSortedArrays(int[] nums1, int[] nums2) {
int m = nums1.length;
int n = nums2.length;
// 奇数长度时 left == right;偶数长度时 left、right 是中间两个位置
int left = (m + n + 1) / 2;
int right = (m + n + 2) / 2;
int leftVal = findKth(nums1, 0, nums2, 0, left);
int rightVal = findKth(nums1, 0, nums2, 0, right);
// 先转 long 再相加,避免极端值域下的 int 溢出
return ((long) leftVal + rightVal) / 2.0;
}
/** 在两个有序数组中找第 k 小的元素,k 从 1 开始计数 */
private int findKth(int[] nums1, int i, int[] nums2, int j, int k) {
// 数组 1 已耗尽,第 k 小完全落在数组 2 中,直接下标定位
if (i >= nums1.length) {
return nums2[j + k - 1];
}
// 数组 2 已耗尽,第 k 小完全落在数组 1 中
if (j >= nums2.length) {
return nums1[i + k - 1];
}
// 只要一个名额,取两个头部较小者
if (k == 1) {
return Math.min(nums1[i], nums2[j]);
}
int half = k / 2;
// 第 half 个元素越界时用 MAX_VALUE 占位,表示"这一段不够,可以整段保留"
int midVal1 = (i + half - 1 < nums1.length) ? nums1[i + half - 1] : Integer.MAX_VALUE;
int midVal2 = (j + half - 1 < nums2.length) ? nums2[j + half - 1] : Integer.MAX_VALUE;
if (midVal1 < midVal2) {
// nums1[i .. i+half-1] 这 half 个元素一定不是第 k 小,整段淘汰
return findKth(nums1, i + half, nums2, j, k - half);
} else {
// nums2[j .. j+half-1] 这 half 个元素一定不是第 k 小
return findKth(nums1, i, nums2, j + half, k - half);
}
}
}
关键代码逐行解释
-
int left = (m + n + 1) / 2; int right = (m + n + 2) / 2;把奇偶两种情况合并成一套代码的技巧。N 为奇数时两式结果相同,偶数时恰好是中间两个下标,所以主函数里永远只写一行平均值公式。
-
if (i >= nums1.length) return nums2[j + k - 1];数组 1 淘汰空了,第 k 小只能落在数组 2 的第
j + k - 1位。这里直接下标定位而不继续递归,是重要剪枝;能保证下标合法,是因为 k 始终不超过两数组剩余元素总数。 -
if (k == 1) return Math.min(nums1[i], nums2[j]);递归出口。k == 1 表示只要最小元素,两个数组的头部(前面都已被淘汰,
i、j就是当前最小值位置)里较小者即答案。 -
int midVal1 = (i + half - 1 < nums1.length) ? nums1[i + half - 1] : Integer.MAX_VALUE;取数组 1 的第 half 个元素。越界时用
MAX_VALUE兜底,使它在比较中永远不小于对方,于是分支自然变成"淘汰另一个数组的 half 段",语义完全正确。 -
return findKth(nums1, i + half, nums2, j, k - half);淘汰 nums1 的
[i, i+half-1]这 half 个元素:起点推到i + half,同时k -= half。注意起点是i + half而不是i + half - 1,后者会把被淘汰的元素留在区间里,导致死循环。
时间、空间复杂度
-
时间复杂度:O(log(m + n))
每次递归 k 变为
k - k/2,规模减半,最多递归log k层;k 最大为(m+n+2)/2,故总层数 O(log(m+n));每层只做常数次比较与分支。 -
空间复杂度:O(log(m + n))
递归栈深度为对数级,没有申请额外数组。若把递归改写成
while循环(手动维护i、j、k),空间可进一步降到 O(1)。
易错点
-
中位数下标公式写错 :
(m+n+1)/2和(m+n+2)/2是两个不同的式子,不能简化成(m+n)/2。奇数长度的正确性完全依赖"两式相等"这个性质,改动任何一个都会把奇数情况算错。 -
求平均值时的 int 溢出 :
(long)(leftVal + rightVal)这种写法是先做 int 加法、再把结果转 long,转换发生在溢出之后,防溢出无效。正确写法是((long) leftVal + rightVal) / 2.0。本题元素范围|nums[i]| ≤ 10^6不会触发,但值域一放大就会翻车。 -
淘汰时下标要多走一步 :被淘汰的是
[i, i+half-1],新起点必须是i + half。写成i + half - 1会重复考察同一个元素,造成死循环或死递归。 -
越界时不能提前返回 :
i + half - 1越界只说明"这个数组剩余不足 half 个",对方数组仍可能有元素需要继续找。用Integer.MAX_VALUE占位参加比较,让逻辑自动抛弃另一个数组,才是正确做法。
可复用模板
"两个有序序列求第 k 小"的通用二分裁剪模板:
java
/** 在有序数组 a(从 ia 起)和 b(从 ib 起)中找第 k 小,k 从 1 开始 */
private int kth(int[] a, int ia, int[] b, int ib, int k) {
// 1) 一方耗尽:直接下标定位
if (ia >= a.length) return b[ib + k - 1];
if (ib >= b.length) return a[ia + k - 1];
// 2) 出口:只要最小元素
if (k == 1) return Math.min(a[ia], b[ib]);
// 3) 各取第 k/2 个,越界用 MAX_VALUE 占位
int half = k / 2;
int va = (ia + half - 1 < a.length) ? a[ia + half - 1] : Integer.MAX_VALUE;
int vb = (ib + half - 1 < b.length) ? b[ib + half - 1] : Integer.MAX_VALUE;
// 4) 谁小淘汰谁的前 half 段,k 同步减半
if (va < vb) return kth(a, ia + half, b, ib, k - half);
else return kth(a, ia, b, ib + half, k - half);
}
主函数里配一句:求中位数时 k 取 (N+1)/2 与 (N+2)/2 各调一次,再取平均。
变体提示:
- 求两个有序数组的第 k 大 :把比较条件取反,或转成求"第
m+n-k+1小"。 - 求多个有序数组的第 k 小:改成在多个头部做二分/堆,升级为多路归并 + 二分答案。
- 求两个有序数组交给"第 k 小"之外的问题(如第 k 近的距离):先用本模板定位候选区间,再在同一结构上二次二分。
相似题及区别
- LeetCode 21 合并两个有序链表:同样是双序列归并,但它必须真正产出合并结果,只能 O(m+n),没有"跳着淘汰"的空间。
- LeetCode 88 合并两个有序数组:从后往前双指针原地合并,O(m+n);本题只关心中间位置,所以可以用二分把同侧元素整段丢弃。
- LeetCode 33 搜索旋转排序数组 :同样是"二分 + 有序性判断",但它是单数组内比较
nums[mid]与nums[left]来判定哪半有序;本题是跨两个数组比较第k/2个元素,状态维度不同。 - LeetCode 153 寻找旋转排序数组中的最小值 :单数组二分找拐点,终止条件是
left == right,right = mid不回退一格;本题是双数组递归减 k,终止条件是k == 1,起点推i + half。