二分查找为什么总是写错边界?
刷了六百多道题,二分查找是我到现在还会在边界上栽跟头的题。别的题错了,要么是思路不对,要么是压根没学过;二分查找不一样,思路三句话就能讲清楚,可一落到 while 条件、mid 加不加一这些细节上,十个人有九个第一次写不对。
后来我想明白一件事:二分查找的边界 bug,几乎都出在同一个地方,就是没想清楚每一轮循环里 lo 和 hi 到底代表什么。这个想清楚了,<= 还是 <、mid + 1 还是 mid,全都自动推出来,不用背。
这篇文章就用这一条思路,把二分查找的几种写法串一遍。代码用 Java。
先立一条不变量
写二分查找之前,先在脑子里立条规矩:
每一轮循环开始的时候,如果
target在数组里,它的下标一定落在[lo, hi]这个闭区间里。
初始化 lo = 0, hi = n - 1,整个数组都在范围里,规矩成立。
后面每一轮只干一件事:取中间点 mid,把 mid 那一半明确排除掉,区间缩一半,同时保证排除掉的那一半里绝不可能有 target。这条不破,规矩就一直在。
nums[mid] < target,target只可能在mid右边,左边界推到mid + 1nums[mid] > target,target只可能在mid左边,右边界推到mid - 1
+1 和 -1 不是摆设:mid 这轮已经查过了,得把它踢出区间,不然区间不缩,甚至死循环。
循环什么时候停?区间里没元素了,也就是 lo > hi。到这一步还没找到,就是真没有。
写法一:找 target
最基础的版本,存在就返回下标,不存在返回 -1。
int binarySearch(int[] nums, int target) {
int lo = 0, hi = nums.length - 1; // [lo, hi] 闭区间
while (lo <= hi) { // 区间里还有元素
int mid = lo + (hi - lo) / 2;
if (nums[mid] < target) {
lo = mid + 1;
} else if (nums[mid] > target) {
hi = mid - 1;
} else {
return mid; // 命中,直接返回
}
}
return -1; // 区间空了
}
为什么是 lo <= hi?因为 hi = n - 1,区间是闭的,lo == hi 的时候区间里还有一个元素,不能停。
mid 写成 lo + (hi - lo) / 2 而不是 (lo + hi) / 2,是防 lo + hi 溢出。结果一样,但 lo + hi 可能爆 int。
写法二:找最左位置
数组里可能有重复元素。比如 nums = [1, 3, 3, 3, 5],target = 3,我要的是第一个 3 的下标 1,不是中间的某个 3。
思路还是那条不变量,区别只在 nums[mid] == target 的时候,别急着返回,继续往左收缩。
int lowerBound(int[] nums, int target) {
int lo = 0, hi = nums.length - 1;
while (lo <= hi) {
int mid = lo + (hi - lo) / 2;
if (nums[mid] < target) {
lo = mid + 1;
} else { // nums[mid] >= target
hi = mid - 1; // 往左压,等于 target 也不停
}
}
if (lo >= nums.length || nums[lo] != target) return -1;
return lo;
}
这里把两个分支并成一个:只要 nums[mid] >= target,就把 hi 压到 mid - 1。等于 target 也不停,区间就一路往左缩,直到 hi 掉到第一个 3 的左边,循环才停。
循环结束那一刻 lo = hi + 1,lo 正好停在第一个 >= target 的位置。如果 nums[lo] 恰好等于 target,lo 就是最左位置。
有个坑:target 比所有元素都大时,lo 会一路冲到 n,也就是数组末尾再往右一格。
所以返回前得先查
lo >= nums.length,不然 nums[lo] 直接越界。
写法三:找最右位置
对称的,找最后一个等于 target 的下标。
int upperBoundMinusOne(int[] nums, int target) {
int lo = 0, hi = nums.length - 1;
while (lo <= hi) {
int mid = lo + (hi - lo) / 2;
if (nums[mid] <= target) { // 等于 target 也继续往右
lo = mid + 1;
} else {
hi = mid - 1;
}
}
if (hi < 0 || nums[hi] != target) return -1;
return hi;
}
这次反过来,nums[mid] <= target 时把 lo 推到 mid + 1,等于 target 也不停,区间一路往右缩,直到 lo 越过最后一个 3。
循环结束时 hi = lo - 1,hi 正好停在最后一个 <= target 的位置。注意这里返回的是 hi 不是 lo,关键就在这个 -1:lo 已经越过 target 了,hi 才还停在它身上。
对称的坑:target 比所有元素都小,hi 会一路退到 -1。
返回前查一下 hi < 0。
三份代码,其实只差一行
把三个函数摆一起看,长得几乎一样:
int binarySearch(int[] nums, int target) {
int lo = 0, hi = nums.length - 1;
while (lo <= hi) {
int mid = lo + (hi - lo) / 2;
if (nums[mid] < target) lo = mid + 1;
else if (nums[mid] > target) hi = mid - 1;
else return mid; // 命中就停
}
return -1;
}
int lowerBound(int[] nums, int target) {
int lo = 0, hi = nums.length - 1;
while (lo <= hi) {
int mid = lo + (hi - lo) / 2;
if (nums[mid] < target) lo = mid + 1;
else hi = mid - 1; // 含等于,往左挤
}
if (lo >= nums.length || nums[lo] != target) return -1;
return lo;
}
int upperBoundMinusOne(int[] nums, int target) {
int lo = 0, hi = nums.length - 1;
while (lo <= hi) {
int mid = lo + (hi - lo) / 2;
if (nums[mid] <= target) lo = mid + 1; // 含等于,往右挤
else hi = mid - 1;
}
if (hi < 0 || nums[hi] != target) return -1;
return hi;
}
差别全在 nums[mid] 跟 target 比较的那个分支上:
- 想找到就停:
==时直接return mid - 想往左挤:
>=时hi = mid - 1,返回lo - 想往右挤:
<=时lo = mid + 1,返回hi
记住这一句就够:命中别停,往左挤返回 lo,往右挤返回 hi。
最后说一句
还有一种写法,把区间设成左闭右开 [lo, hi),hi 初始化为 n。Java 的 Arrays.binarySearch、C++ 的 lower_bound 内部都是这么干的,半开区间在表示"空区间"和"插入位置"上更顺手。
但我自己写还是习惯闭区间。闭区间里 lo 和 hi 的对称性一眼能看出来,+1、-1、<= 这些细节靠不变量都能推出来,不容易错。选哪种都行,关键是写之前把不变量说清楚,别背。