二分查找为什么总是写错边界?手把手教你

0 阅读5分钟

二分查找为什么总是写错边界?

刷了六百多道题,二分查找是我到现在还会在边界上栽跟头的题。别的题错了,要么是思路不对,要么是压根没学过;二分查找不一样,思路三句话就能讲清楚,可一落到 while 条件、mid 加不加一这些细节上,十个人有九个第一次写不对。

后来我想明白一件事:二分查找的边界 bug,几乎都出在同一个地方,就是没想清楚每一轮循环里 lohi 到底代表什么。这个想清楚了,<= 还是 <mid + 1 还是 mid,全都自动推出来,不用背。

这篇文章就用这一条思路,把二分查找的几种写法串一遍。代码用 Java。

先立一条不变量

写二分查找之前,先在脑子里立条规矩:

每一轮循环开始的时候,如果 target 在数组里,它的下标一定落在 [lo, hi] 这个闭区间里。

初始化 lo = 0, hi = n - 1,整个数组都在范围里,规矩成立。

后面每一轮只干一件事:取中间点 mid,把 mid 那一半明确排除掉,区间缩一半,同时保证排除掉的那一半里绝不可能有 target。这条不破,规矩就一直在。

  • nums[mid] < targettarget 只可能在 mid 右边,左边界推到 mid + 1
  • nums[mid] > targettarget 只可能在 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。

bs-leftmost.png

思路还是那条不变量,区别只在 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 + 1lo 正好停在第一个 >= target 的位置。如果 nums[lo] 恰好等于 target,lo 就是最左位置。

有个坑:target 比所有元素都大时,lo 会一路冲到 n,也就是数组末尾再往右一格。

bs-leftmost-oob.png 所以返回前得先查 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;
}

bs-rightmost.png

这次反过来,nums[mid] <= target 时把 lo 推到 mid + 1,等于 target 也不停,区间一路往右缩,直到 lo 越过最后一个 3。

循环结束时 hi = lo - 1hi 正好停在最后一个 <= target 的位置。注意这里返回的是 hi 不是 lo,关键就在这个 -1lo 已经越过 target 了,hi 才还停在它身上。

对称的坑:target 比所有元素都小,hi 会一路退到 -1。

bs-rightmost-oob.png

返回前查一下 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 内部都是这么干的,半开区间在表示"空区间"和"插入位置"上更顺手。

但我自己写还是习惯闭区间。闭区间里 lohi 的对称性一眼能看出来,+1-1<= 这些细节靠不变量都能推出来,不容易错。选哪种都行,关键是写之前把不变量说清楚,别背