回溯法入门:从递归穷举到剪枝优化的完整指南
“初识回溯法”这个标题看起来像是学习笔记又像是一篇技术分享的题目。作为一个跟算法打了多年交道的人我想借这个题目把我对回溯法的理解、踩过的坑、总结的套路原原本本写出来。这篇文章不是什么高深的理论手册而是从一个写代码的人的角度出发讲清楚回溯法是什么、怎么想、怎么写、怎么优化。如果你正在学算法、刷题或者准备面试这篇文章应该能帮你省下不少走弯路的时间。1. 回溯法到底是个什么东西1.1 从直觉出发它其实就是带后悔药的穷举很多人一听到“回溯”两个字就觉得是个很高深的东西。其实完全不是。回溯法的本质就是穷举。只不过它这个穷举是带着“后悔药”的穷举。举个例子。你走一个迷宫到了一个岔路口先往左边走。走了几步发现是死胡同你会怎么做当然是退回来换右边那条路继续走。这个“退回来”的动作就是回溯。而走迷宫这个过程本身就是在把所有可能的路都试一遍试到能走出迷宫的那条为止——这就是穷举。再比如你玩数独。一个格子能填好几个数字你先填了一个往下填发现矛盾了于是你把填的数字擦掉换另一个数字继续试。这个“擦掉重填”的动作也是回溯。所以回溯法一点都不神秘。它就是“尝试-失败-撤销-再尝试”的循环。程序里的回溯只是把迷宫和数独里的这个动作变成了一套可以复用的代码逻辑。1.2 回溯法的三个关键点路径、选择、撤销理解回溯法我建议你先记住三个词路径、选择列表、结束条件。路径就是你到当前这一步为止已经做出的所有选择。比如全排列问题里你已经在第一位填了1、第二位填了2那“1,2”就是你当前的路径。选择列表在当前状态下你还可以做哪些选择。上面例子里如果数字总共是1、2、3那你接下来只能选3了。结束条件什么时候算找到了一个答案或者什么时候继续下去已经没有意义了。比如全排列的路径长度已经等于数组长度这时候就找到一个完整排列了可以记录结果并返回。这三个词不是我的发明它的思想跟递归回溯框架完全对应。你抓住这三个关键点后面看所有回溯题都会觉得眼熟。而回溯法的代码之所以叫“回溯”就是因为它在递归返回的时候需要做撤销这一步。这个撤销操作是整个回溯框架的灵魂也是新手最容易出错的地方。说实话我见过太多同学代码思路是对的就是因为忘记撤销结果结果集里全是同一个路径的重复复制排查半天找不出原因。1.3 它解决的是什么问题枚举型问题的通用解法那我为什么要用回溯法或者说什么问题该用回溯法我的经验是当你遇到一个问题它的解空间是“有限但巨大”的而且你需要在里面找到满足某些条件的解时回溯法就是一个非常自然的解法。举几个典型的例子排列类给一个数组求所有排列。组合类从n个数里选k个求所有组合。子集类给一个数组求所有子集。搜索类八皇后、数独、迷宫寻路、单词搜索。表达式类给一串数字插入运算符使结果等于目标值。这类问题的共同特点是解空间是一棵树或图你要遍历这棵树的所有分支但又不是傻傻地把每个分支都走到底——遇到明显不行的分支你要提前停下来这就是剪枝走不通的分支你要退回去换一条路这就是回溯。如果用传统多层循环去写排列3个数得写3层循环排列10个数就得写10层循环这显然不现实。而回溯法用递归来动态地控制循环层数不管排列多少个元素代码框架都是一样的。这一点是回溯法相比暴力写法的最大优势。2. 回溯法的通用框架与代码模板2.1 先看一个标准的回溯模板直接说结论回溯法有一个非常固定的模板几乎所有回溯题都是在这个模板上做修改。你先把这个模板背下来再慢慢理解它。Python版本的模板大概是这样的def backtrack(路径, 选择列表): if 满足结束条件: 记录结果 return for 选择 in 选择列表: 做出选择 backtrack(路径, 选择列表) 撤销选择就这么点东西。所有的回溯题都是在这几行代码的“骨架子”上加肉。我再用一个更具体、更接近实际题目的模板来展示。以“全排列”为例def permute(nums): result [] path [] used [False] * len(nums) def backtrack(): # 结束条件路径长度等于数组长度说明所有元素都用完了 if len(path) len(nums): result.append(path[:]) # 注意要复制一份path return for i in range(len(nums)): if used[i]: continue # 做出选择 used[i] True path.append(nums[i]) backtrack() # 撤销选择 used[i] False path.pop() backtrack() return result注意看做出选择和撤销选择是对称的。你在递归之前做了什么递归返回之后就必须把那一步做的事情恢复原样。这个对称性是回溯代码正确性的关键。如果破坏了对称性程序就会出错——要么结果不对要么栈溢出。2.2 模板里每一行到底在干什么初学者最容易犯的毛病是背下来了模板但不知道每一行为什么要这样写。我来逐行拆解一下。结束条件判断if len(path) len(nums): result.append(path[:]) return这里的目的是告诉递归什么时候该停止。在全排列里路径长度达到了数组长度说明我们已经选了所有数字这就是一个合法排列。此时把路径复制一份放进结果集。注意有个细节我写的是path[:]不是path。这是Python里一个特别经典的坑。因为path是一个引用如果你直接append(path)后面回溯的时候 path 一变化结果集里已经存下的东西也会跟着变。最后你会发现结果集里的所有排列全都长得一模一样。这一点一定要养成肌肉记忆存结果时永远要拷贝一份快照。循环遍历选择列表for i in range(len(nums)):这一步是在枚举“当前状态下所有可能的下一步选择”。每轮循环就是在尝试一种可能性。for循环跑一遍就相当于把当前节点下的所有子节点都访问了一遍。跳过已使用的元素if used[i]: continue这是用来维护“选择列表”的。因为排列问题里每个数字只能用一次。已经用过的数字就不能再选了。所以在进入递归前需要检查一下这个元素是否已经在路径中了。这里我要补充一句不同的题目维护“选择列表”的方式完全不同。有的用used数组比如排列有的用start起始索引比如组合和子集有的干脆用集合。后面讲具体题目的时候你会体会到这种差异。做出选择与撤销选择used[i] True path.append(nums[i]) backtrack() used[i] False path.pop()这里的逻辑可以理解为我先选了nums[i]然后往下一层递归去处理剩下的选择。递归返回之后说明“以 nums[i] 开头的情况”已经全部处理完了这时候我必须把状态恢复成“还没有选 nums[i]”的样子再去尝试nums[i1]。如果我把撤销选择那两行注释掉会怎么样程序会一直往深处递归直到path的长度超过nums的长度然后越陷越深最终栈溢出。这就是对称性被破坏的后果。2.3 画递归树理解回溯最快的方法讲了这么多概念我还是要说理解回溯法最快的方式是手动画递归树。拿[1,2,3]全排列举例。初始是空路径然后第一层选1、选2、选3分三个分支。第二层以“选了1”为例剩下的选择是2和3于是又分出两个分支。第三层以“选了1、2”为例剩下的选择只剩3没得选直接到达叶子节点。画出来是一棵三层的树长得非常规整。我跟你说一个当时把我点醒的话回溯法其实就是对这棵树的深度优先遍历。for循环就是在访问当前节点的所有子节点递归就是在进入子节点递归返回就是在从子节点回到父节点。而“撤销选择”这个动作正是在DFS中从一个分支返回另一个分支时所必需的状态清理。你一旦建立起“递归树”的视角看任何回溯题都会变得特别清晰。遇到一道题先在纸上把这棵树的形状画出来搞清楚每个节点代表什么状态、每个分支代表什么选择代码其实就是把这棵树用DFS走一遍。3. 经典入门题拆解全排列、组合、子集3.1 全排列最典型的回溯入门题前面已经给出了全排列的核心代码这里我完整展开一下并给你演示它的执行过程。def permute(nums): result [] path [] used [False] * len(nums) def backtrack(): if len(path) len(nums): result.append(path[:]) return for i in range(len(nums)): if used[i]: continue used[i] True path.append(nums[i]) backtrack() used[i] False path.pop() backtrack() return result以nums [1, 2, 3]为例backtrack()开始时 path 为空。i0选1path[1]。进入递归。i0used[0]True跳过。i1选2path[1,2]。进入递归。i0、1 都被占用。i2选3path[1,2,3]。进入递归。path长度3等于nums长度把 [1,2,3] 存进result返回。撤销选3path[1,2]。返回上一层。i2撤销选2path[1]。返回上一层。i1撤销选1path[]。i1选2path[2]……接下来会得到 [2,1,3]、[2,3,1]。整个过程中你观察path的变化会发现它就像一个栈不断压入新元素、弹出旧元素。这就是回溯最直观的样子。补充一个复杂度说明全排列的总数是 n! 个而每个排列复制到结果里需要 O(n) 的时间所以总时间复杂度是 O(n×n!)。空间复杂度主要是递归栈深度为 O(n)加上结果集占用的 O(n×n!)。3.2 组合问题从全排列到组合的差异组合问题跟全排列长得很像但有一个重要区别组合不关心顺序。[1,2]和[2,1]在排列里是两个不同结果但在组合里只算一个。以“从1到n中选k个数的所有组合”为例就是在各大刷题网站上那道高频题。如果用全排列的思路去写会产生大量重复结果。怎么办答案是引入一个start 参数让选择只能在当前数字之后进行这样就天然避免了顺序不同的重复。来看代码def combine(n, k): result [] path [] def backtrack(start): # 结束条件路径里已经选了k个数 if len(path) k: result.append(path[:]) return # 从start开始选择避免重复组合 for i in range(start, n 1): path.append(i) backtrack(i 1) path.pop() backtrack(1) return result注意区别这里没有used数组因为start已经保证了一个数字不会被重复使用而且后一个数字一定比前一个数字大因此不会出现[2,1]这种重复组合。这里有个小优化也是面试官很爱问的点剪枝。如果当前路径长度是len(path)还需要选k - len(path)个数而剩余可选数字最多到 n那么i的最大取值应该是n - (k - len(path)) 1。如果i再大剩下的数字不够凑满 k 个这一层递归注定徒劳。所以循环可以写成for i in range(start, n - (k - len(path)) 2):这个优化在数据量大的时候效果很明显。比如 n100、k50剪枝能砍掉大量无效分支。3.3 子集问题换一种视角看回溯子集问题更加简单但思路跟排列、组合略有不同它要求的是所有子集而不只是固定长度的组合。也就是说路径在任何长度下都可以作为答案被收集。以“给定一个不含重复元素的数组 nums返回所有子集”为例def subsets(nums): result [] path [] def backtrack(start): # 每个状态的path都是一个合法的子集都记录下来 result.append(path[:]) for i in range(start, len(nums)): path.append(nums[i]) backtrack(i 1) path.pop() backtrack(0) return result仔细看这个代码和组合问题的代码只有一处不同收集结果的时机。组合问题是等到 path 长度等于 k 才收集子集问题是每进入一层就收集。因为子集不要求长度任何长度都是合法答案。你也可以换一种写法用“选或不选”的思路def subsets(nums): result [] path [] def backtrack(index, path): if index len(nums): result.append(path) return # 不选当前元素 backtrack(index 1, path) # 选当前元素 backtrack(index 1, path [nums[index]]) backtrack(0, []) return result这种写法本质上是对数组的每一位做“选/不选”的二叉决策形成的是一棵二叉树。两种思路都能得到正确答案但第一种写法用了start参数和组合、排列的模板更统一我个人推荐你重点掌握第一种。子集问题的结果是 2^n 个复制每个结果需要 O(n) 时间所以时间复杂度 O(n×2^n)空间复杂度 O(n)。这类指数级复杂度在常规数据量下勉强可用但你要心里有数只要回溯法跑起来数据规模稍微大一点时间开销就压不住了。4. 剪枝优化省掉那些注定走不通的路4.1 为什么必须聊剪枝回溯法说白了就是穷举穷举的代价是巨大的。但你可以通过剪枝让这种穷举变得聪明一点。什么叫剪枝想象你在遍历一棵决策树树上有些分支你还没走下去就已经知道它不可能产生合法解了。那你还傻傻地走吗当然不。直接跳过这就是剪枝。剪枝不是说它能让“不可能变可能”而是说它能把“无用功”提前止损。它不改变算法的上界复杂度但在实际运行中经常能把时间从几小时降到几秒。很多人写回溯题超时不是思路错了而是少了剪枝这一步。4.2 常见剪枝策略我总结了一下常见的剪枝策略大致有这几种。第一种排列组合中的“剩余不足”剪枝。这个我在组合问题里已经提到过就是判断“当前已选元素 剩余可选元素”是否还够凑成目标长度。不够就直接返回。if len(path) (n - start 1) k: return第二种排序 同层去重剪枝。这种主要用在数组里有重复元素的题目里比如“组合总和 II”“全排列 II”。思路是先对数组排序让相同的元素相邻。然后在一个循环里如果发现当前元素和前一个元素相同并且前一个元素还没有被使用过那说明这个分支在上一个循环中已经处理过了直接跳过。if i 0 and nums[i] nums[i - 1] and not used[i - 1]: continue这一行代码看起来不起眼实际上能避免大量重复解的产生。第三种可行性剪枝。这个比较灵活要针对具体题目的约束条件来设计。比如八皇后问题里如果你已经能判断当前位置和已有皇后冲突了那就没必要递归下去了。再比如一些方块填充问题如果当前方案已经打破约束直接剪掉。4.3 剪枝实战八皇后问题完整拆解聊了这么多理论我拿八皇后问题给你完整演示一下“回溯 剪枝”是怎么配合的。问题描述在 n×n 的棋盘上放置 n 个皇后使得任意两个皇后不能出现在同一行、同一列、同一对角线上。返回所有合法摆法。思路逐行放置皇后。因为每行只能放一个皇后当我处理第 row 行时只需要考虑在每一列能不能放。判断条件有三个该列没有被占用。左上到右下的对角线没有被占用。右上到左下的对角线没有被占用。这里有一个经典技巧对于棋盘上的坐标 (r, c)左上到右下的对角线可以用r - c来标记值是常数右上到左下的对角线可以用r c来标记值也是常数。用两个集合来记录这些对角线是否被占用判断就变成了 O(1) 的哈希查找。def solve_n_queens(n): result [] board [[.] * n for _ in range(n)] col_used set() diag1_used set() # r - c 恒定 diag2_used set() # r c 恒定 def backtrack(row): if row n: # 收集当前棋盘状态 result.append([.join(r) for r in board]) return for col in range(n): if col in col_used or (row - col) in diag1_used or (row col) in diag2_used: continue # 不可行剪枝 # 做选择 board[row][col] Q col_used.add(col) diag1_used.add(row - col) diag2_used.add(row col) backtrack(row 1) # 撤销选择 board[row][col] . col_used.discard(col) diag1_used.discard(row - col) diag2_used.discard(row col) backtrack(0) return result注意观察每行只放一个皇后这本身就是一种隐式的剪枝——我们从来不会在同一行尝试放两个皇后。列冲突、对角线冲突的判断则是在每次尝试前就提前止损。这就是“可行性剪枝”最典型的应用。八皇后问题如果不做任何剪枝需要检查 n^n 种摆法n8 时是 1600 多万种。加了剪枝之后实际递归的节点数会大幅下降n8 时合法摆法只有 92 种而整个搜索空间中被访问的有效节点也就几千个。这差距就是剪枝的力量。5. 常见问题与调试技巧5.1 新手最常踩的坑我见过的回溯代码报错来来回回就那么几个原因。我把它们整理成一张表你遇到问题可以对照排查。常见问题表现原因解决方法结果集全是同一个路径输出很多一模一样的列表存储结果时没有拷贝保存的是同一个引用的快照用path[:]或path.copy()存结果缺少大量结果答案数量明显偏少撤销选择没做或撤销位置不对导致状态污染确保递归前后状态完全对称结果重复出现了多个相同的排列/组合没有使用start参数或没有对重复元素去重组合问题用start重复元素场景先排序后去重死循环/栈溢出程序一直跑不停或直接报递归深度超限结束条件写错或递归参数没有向终止条件逼近检查递归出口打印路径观察递归深度变量作用域导致的结果错乱结果集直接报错或值不对在递归里改了不可变对象或误用了可变默认参数不要在递归参数里用列表做默认值改用闭包或显式传递我特别要强调第一行也就是“存储结果时没有拷贝”的问题。这是所有回溯新手都会踩的坑没有例外。哪怕是你已经理解了回溯逻辑也容易在写的时候一顺手就result.append(path)。等到打印结果才发现全是同样的列表。我当时的习惯是凡是往结果集里放数据的地方一律写path[:]并且写完代码之后回头检查一遍每个 append 的地方。5.2 调试回溯代码的有效方法回溯代码有个特点它不像普通的顺序逻辑你很难用断点直观地看结果。因为递归一层层嵌套断点太多根本看不过来。我常用的手段有这么几个。第一打印路径法。在递归函数开头加一行 print观察每次进入递归时的 path 是什么。这样你就能在控制台看到整个回溯过程的轨迹哪里多走了、哪里少走了一目了然。print(f进入递归当前路径: {path})第二缩小输入规模法。如果你在 n10 的时候跑出问题根本不知道错在哪。先把它改成 n3手动在纸上把正确结果写出来再对比程序输出。一旦 n3 的输出和你手写的结果一致再逐步放大。这个习惯能救你无数次。第三画递归树对比法。程序输出不对的时候我通常会在纸上先把递归树画出来标出哪些路径是合法的哪些是被剪掉的然后对照代码看代码的每一步在树上对应什么位置。这个方法看起来笨但效率极高。尤其是当你对某个分支的跳转逻辑产生怀疑时画树比单纯盯代码管用得多。5.3 关于时间和空间复杂度心里要有个底回溯法的时间复杂度很多新手喜欢直接问“是多少”。但这个问题真的没有统一答案因为它完全取决于解空间的大小和剪枝的强度。全排列解空间是 n!时间复杂度 O(n×n!)。子集解空间是 2^n时间复杂度 O(n×2^n)。组合解空间是 C(n, k)时间复杂度 O(k×C(n, k))。八皇后无剪枝是 O(n^n)有冲突剪枝后会大幅下降但上界仍是 O(n!)。空间复杂度相对好判断一些主要的开销来自递归栈的深度以及保存结果所需的空间。递归栈深度一般等于答案长度即 O(n)。结果集的空间取决于合法解的数量这个只能根据具体问题估算。我想提醒你的是如果你发现回溯代码在 n30 以上的数据量级还能秒出结果那大概率是剪枝发挥了奇效而不是回溯本身变快了。回溯法本质上是暴力搜索不要指望它在指数级问题上有什么“银弹”。真正有用的是让该剪的剪掉、该停的停下用最小代价完成搜索。前面也提到过当你对回溯的理解深了之后你会发现它跟树的深度优先遍历、跟动态规划的状态转移、甚至跟图论里的拓扑排序在思想上有千丝万缕的联系。但这些是后面的话题了。回到开头我在想的问题为什么那么多人觉得回溯难后来我想明白了难的不是模板而是你还没建立“递归树”的直觉。一旦你想通了每个节点就是一层状态、每条边就是一个选择、每个叶子就是一个答案回溯法的大门就算是真正推开了。最后再分享一个小技巧刷题的时候把全排列、组合、子集这三个问题放在一起研究。它们长得像思路互相关联把它们彻底吃透你对回溯法的理解会比单做十道难题都扎实。