还没走到底,就提前发现这条路走不通——直接放弃,省下后面一大段本该白费的功夫。
12.2 节的回溯,每一层都会把"还没用过的数"完整地试一遍——如果问题规模变大,分支数量会爆炸性增长(比如把 8 个数排列,就有 8! = 40320 种可能)。剪枝要解决的问题是:很多分支其实还没走到底就已经能看出肯定不行了,与其老老实实递归到最深处才发现白费功夫,不如提前判断出来,直接放弃这整条分支,连递归都不用往下走。
本节用经典的 N 皇后问题来演示剪枝:在 N×N 的棋盘上放 N 个皇后,要求任意两个皇后都不能在同一行、同一列、同一条对角线上。为了方便,可以规定"每一行必须且只能放一个皇后"——这样问题就变成了:给第 0 行到第 N-1 行分别分配一个列号,列号不能重复。这其实和 12.2 节的"排列"是同一件事:给 N 行分配 0~N-1 这 N 个列号,就是这 N 个数的一种排列,只是这次多了一条"不能在同一条对角线上"的额外规则。
最直接的想法:完全照搬 12.2 节的全排列代码,先把 N 行的所有排列方式都生成出来(列号不重复这条规则,在生成排列的过程中天然就满足了),每生成一个完整的排列,才检查它有没有任意两个皇后落在同一条对角线上。
[0,1,2,3] 表示第 0 行放第 0 列、第 1 行放第 1 列……结果四个皇后正好落在同一条对角线上,是明显不合法的方案。但基础写法在生成这个排列的过程中,其实第 1 行放 1 的时候就已经能看出和第 0 行冲突了——只是代码里没有做这个判断,非要等 4 行全部放完才检查。| 1 | int n = 4; |
| 2 | int path[20]; // path[i] 表示第 i 行的皇后放在哪一列 |
| 3 | bool used[20]; |
| 4 | int count = 0; // 合法方案的总数 |
| 5 | |
| 6 | bool IsValid() // 检查 path 里已经排满的这一整个排列,有没有对角线冲突 |
| 7 | { |
| 8 | for (int i = 0; i < n; i++) |
| 9 | for (int j = i + 1; j < n; j++) // 两两比较所有皇后 |
| 10 | if (abs(path[i] - path[j]) == abs(i - j)) // 行差 == 列差,说明在同一条对角线 |
| 11 | return false; |
| 12 | return true; |
| 13 | } |
| 14 | |
| 15 | void Backtrack(int depth) |
| 16 | { |
| 17 | if (depth == n) // 排列生成完了,这时候才检查 |
| 18 | { |
| 19 | if (IsValid()) { count++; } |
| 20 | return; |
| 21 | } |
| 22 | for (int i = 0; i < n; i++) |
| 23 | { |
| 24 | if (used[i]) { continue; } |
| 25 | path[depth] = i; used[i] = true; |
| 26 | Backtrack(depth + 1); |
| 27 | used[i] = false; |
| 28 | } |
| 29 | } |
depth == n 意味着——不管中间过程冲突得多明显,代码都会老老实实地把这一整条分支递归到最深处,才在第 19 行检查一次。以 [0,1,……] 开头的排列一共有 (n-2)! 种(剩下的行随便排),而这些排列全部都会因为第 0 行和第 1 行本身就冲突而不合法——但代码要把它们一个不漏地全部生成出来,才能发现这一点。与其等排列生成完再检查,不如每往下放一个新皇后,就立刻检查它和已经放置的所有皇后是否冲突——只要发现冲突,直接放弃这个位置,连递归都不用调用,更不用说把剩下的行也走一遍了。
1,列差也是 1,正好在同一条对角线上,立刻就能判断出冲突(红色 ✕)。剪枝之后,代码根本不会递归进入第 2 行,直接放弃这个选择、换第 1 行的下一个列继续试——而基础写法要把第 2、3 行的所有排列方式都生成完(2! = 2 种)才能发现同样的结论。| 1 | bool Conflict(int row, int col) // 判断新皇后 (row,col) 和已放置的皇后是否冲突 |
| 2 | { |
| 3 | for (int i = 0; i < row; i++) // 只需要看前面已经放好的这些行 |
| 4 | { |
| 5 | if (abs(path[i] - col) == abs(i - row)) // 行差 == 列差,在同一条对角线 |
| 6 | { return true; } |
| 7 | } |
| 8 | return false; |
| 9 | } |
| 10 | |
| 11 | void Backtrack(int depth) |
| 12 | { |
| 13 | if (depth == n) { count++; return; } // 能走到这里,说明前面全部合法,不用再检查了 |
| 14 | for (int i = 0; i < n; i++) |
| 15 | { |
| 16 | if (used[i]) { continue; } |
| 17 | if (Conflict(depth, i)) { continue; } // ★ 剪枝:这里冲突就直接跳过,连递归都不调用 |
| 18 | path[depth] = i; used[i] = true; |
| 19 | Backtrack(depth + 1); |
| 20 | used[i] = false; |
| 21 | } |
| 22 | } |
优化一虽然提前剪了枝,但 Conflict 函数每次都要遍历一遍前面所有已经放置的皇后(第 3 行的 for 循环),判断一次的开销是 O(depth)。可以换个思路:用几个数组分别记录"这一列""这条对角线"是否已经被占用,放置或撤销的时候同步更新这些数组,判断就只需要查一次数组,变成 O(1)。
| 数组 | 下标怎么算 | 含义 |
|---|---|---|
colUsed[col] | 直接用列号 | 第 col 列是否已经有皇后 |
diag1[row - col + n] | 行号减列号(加 n 避免负数下标) | 同一条"左上到右下"对角线上,row - col 是常数 |
diag2[row + col] | 行号加列号 | 同一条"右上到左下"对角线上,row + col 是常数 |
| 1 | bool colUsed[20], diag1[40], diag2[40]; // 对角线最多有 2n-1 条,数组要开够 |
| 2 | |
| 3 | void Backtrack(int depth) |
| 4 | { |
| 5 | if (depth == n) { count++; return; } |
| 6 | for (int col = 0; col < n; col++) |
| 7 | { |
| 8 | if (colUsed[col] || diag1[depth - col + n] || diag2[depth + col]) |
| 9 | { continue; } // ★ 三个数组一查就知道冲不冲突,不用遍历 |
| 10 | |
| 11 | colUsed[col] = diag1[depth - col + n] = diag2[depth + col] = true; // 三个数组同时标记占用 |
| 12 | Backtrack(depth + 1); |
| 13 | colUsed[col] = diag1[depth - col + n] = diag2[depth + col] = false; // 三个数组同时撤销 |
| 14 | } |
| 15 | } |
colUsed、diag1、diag2 三个数组,第 13 行也必须把这三个数组同时改回 false,少改一个都会让后面的分支读到不该存在的"占用"记录,把明明合法的位置误判成冲突。三种写法的核心区别:
| 写法 | 什么时候发现冲突 | 单次判断开销 |
|---|---|---|
| ② 基础写法 | 排列生成完(第 n 行)才检查 | O(n²)(两两比较所有皇后) |
| ③ 优化一 | 每放一个新皇后就检查,冲突立刻停止 | O(depth)(遍历前面已放置的皇后) |
| ④ 优化二 | 同优化一,提前发现 | O(1)(查三个数组) |
n 变大、递归层数很深时,这个常数优化累积起来也很可观。两种优化经常一起使用:用剪枝减少要走的分支数量,再用更快的判断方式降低每个分支的开销。continue,压根不调用递归;如果已经调用了递归、深入到子树内部才判断退出,那只是普通的递归终止条件,并没有省下"生成这棵子树"的开销。剪枝的关键在于判断必须发生在递归调用之前。colUsed、diag1、diag2 必须同步标记、同步撤销。实际写代码时如果拆成三行分别写,很容易漏掉其中一行的撤销,导致状态"越攒越脏",后面的分支会莫名其妙地被判断为冲突(或者反过来,本该冲突的却被判断为合法)。depth - col 的取值范围是 -(n-1) 到 n-1,直接当数组下标会出现负数、导致越界访问。必须像第 8、11、13 行那样统一加上 n,把整个范围平移到非负区间。