递归与回溯
00:00
递归思想、回溯算法框架、经典回溯问题(子集、排列、组合、N皇后)与剪枝优化。
1. 递归基础
1.1 递归的本质
递归是函数直接或间接调用自身的编程技巧。每个递归必须包含两个要素:
- 基线条件(Base Case):递归终止的条件,防止无限递归
- 递归条件(Recursive Case):将问题分解为更小的子问题
# 经典递归:阶乘
def factorial(n):
if n <= 1: # 基线条件
return 1
return n * factorial(n - 1) # 递归条件
# 经典递归:斐波那契数列
def fibonacci(n):
if n <= 1:
return n
return fibonacci(n - 1) + fibonacci(n - 2)
1.2 递归的执行过程
递归的执行分为两个阶段:
- 递推阶段:不断分解问题,直到达到基线条件
- 回归阶段:从基线条件开始,逐层返回结果
factorial(4) 的执行过程:
递推:f(4) → 4*f(3) → 4*3*f(2) → 4*3*2*f(1) → 4*3*2*1
回归:4*3*2*1 = 24
1.3 递归与迭代的转换
任何递归都可以转换为迭代,但递归通常更直观:
# 递归版本
def sum_recursive(arr, n):
if n == 0:
return 0
return arr[n - 1] + sum_recursive(arr, n - 1)
# 迭代版本
def sum_iterative(arr):
total = 0
for num in arr:
total += num
return total
1.4 递归的复杂度分析
| 递归类型 | 时间复杂度 | 空间复杂度 | 示例 |
|---|---|---|---|
| 线性递归 | O(n) | O(n) | 阶乘、数组求和 |
| 二分递归 | O(2^n) | O(n) | 斐波那契(无记忆化) |
| 尾递归 | O(n) | O(1)* | 尾递归优化的阶乘 |
| 树形递归 | O(b^d) | O(d) | 回溯搜索 |
*尾递归优化需要编译器/解释器支持,Python不支持尾递归优化
2. 回溯算法框架
2.1 回溯算法的核心思想
回溯算法是一种通过探索所有可能候选解来找出所有解的算法。如果候选解被确认不是一个解(或者至少不是最后一个解),回溯算法会通过在上一步进行一些变化来丢弃该解,即”回溯”。
回溯 = 深度优先搜索 + 剪枝 + 状态重置
2.2 回溯算法模板
def backtrack(路径, 选择列表):
if 满足结束条件:
结果.append(路径.copy()) # 注意要拷贝
return
for 选择 in 选择列表:
# 做选择
路径.append(选择)
# 递归
backtrack(路径, 新的选择列表)
# 撤销选择(回溯)
路径.pop()
2.3 回溯的三种典型问题
| 问题类型 | 特征 | 示例 |
|---|---|---|
| 子集问题 | 从集合中选取元素 | 子集、子集II |
| 组合问题 | 从集合中选k个元素 | 组合总和、电话号码 |
| 排列问题 | 元素的排列顺序 | 全排列、字符串排列 |
3. 经典回溯问题
3.1 子集问题
def subsets(nums):
"""LeetCode 78: 子集"""
result = []
def backtrack(start, path):
result.append(path[:]) # 收集所有节点(不仅是叶子节点)
for i in range(start, len(nums)):
path.append(nums[i])
backtrack(i + 1, path) # i+1 避免重复选取
path.pop()
backtrack(0, [])
return result
# 含重复元素的子集
def subsets_with_dup(nums):
"""LeetCode 90: 子集II - 含重复元素"""
nums.sort() # 排序使重复元素相邻
result = []
def backtrack(start, path):
result.append(path[:])
for i in range(start, len(nums)):
# 剪枝:跳过重复元素
if i > start and nums[i] == nums[i - 1]:
continue
path.append(nums[i])
backtrack(i + 1, path)
path.pop()
backtrack(0, [])
return result
# 示例
print(subsets([1, 2, 3]))
# [[], [1], [1,2], [1,2,3], [1,3], [2], [2,3], [3]]
print(subsets_with_dup([1, 2, 2]))
# [[], [1], [1,2], [1,2,2], [2], [2,2]]
3.2 组合问题
def combine(n, k):
"""LeetCode 77: 组合 - 从1..n中选k个数"""
result = []
def backtrack(start, path):
if len(path) == k:
result.append(path[:])
return
# 剪枝:剩余元素不足时提前终止
# 还需要 k - len(path) 个元素,剩余可选范围为 [start, n]
# 所以 i 的上界为 n - (k - len(path)) + 1
for i in range(start, n - (k - len(path)) + 2):
path.append(i)
backtrack(i + 1, path)
path.pop()
backtrack(1, [])
return result
def combination_sum(candidates, target):
"""LeetCode 39: 组合总和 - 可重复选取"""
result = []
def backtrack(start, path, remaining):
if remaining == 0:
result.append(path[:])
return
if remaining < 0:
return
for i in range(start, len(candidates)):
path.append(candidates[i])
backtrack(i, path, remaining - candidates[i]) # 注意是i不是i+1,允许重复
path.pop()
backtrack(0, [], target)
return result
# 示例
print(combine(4, 2))
# [[1,2], [1,3], [1,4], [2,3], [2,4], [3,4]]
print(combination_sum([2, 3, 6, 7], 7))
# [[2,2,3], [7]]
3.3 排列问题
def permute(nums):
"""LeetCode 46: 全排列"""
result = []
def backtrack(path, used):
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(path, used)
path.pop()
used[i] = False
backtrack([], [False] * len(nums))
return result
def permute_unique(nums):
"""LeetCode 47: 全排列II - 含重复元素"""
nums.sort()
result = []
def backtrack(path, used):
if len(path) == len(nums):
result.append(path[:])
return
for i in range(len(nums)):
if used[i]:
continue
# 剪枝:同一层中跳过重复元素
if i > 0 and nums[i] == nums[i - 1] and not used[i - 1]:
continue
used[i] = True
path.append(nums[i])
backtrack(path, used)
path.pop()
used[i] = False
backtrack([], [False] * len(nums))
return result
# 示例
print(permute([1, 2, 3]))
# [[1,2,3], [1,3,2], [2,1,3], [2,3,1], [3,1,2], [3,2,1]]
print(permute_unique([1, 1, 2]))
# [[1,1,2], [1,2,1], [2,1,1]]
3.4 N皇后问题
def solve_n_queens(n):
"""LeetCode 51: N皇后"""
result = []
def is_valid(board, row, col):
# 检查同列
for i in range(row):
if board[i][col] == 'Q':
return False
# 检查左上对角线
i, j = row - 1, col - 1
while i >= 0 and j >= 0:
if board[i][j] == 'Q':
return False
i -= 1
j -= 1
# 检查右上对角线
i, j = row - 1, col + 1
while i >= 0 and j < n:
if board[i][j] == 'Q':
return False
i -= 1
j += 1
return True
def backtrack(row, board):
if row == n:
result.append([''.join(r) for r in board])
return
for col in range(n):
if not is_valid(board, row, col):
continue
board[row][col] = 'Q'
backtrack(row + 1, board)
board[row][col] = '.'
board = [['.' for _ in range(n)] for _ in range(n)]
backtrack(0, board)
return result
# 示例
solutions = solve_n_queens(4)
for sol in solutions:
for row in sol:
print(row)
print()
4. 剪枝优化
4.1 剪枝策略分类
| 剪枝类型 | 说明 | 示例 |
|---|---|---|
| 排序剪枝 | 排序后跳过重复元素 | 子集II、排列II |
| 边界剪枝 | 提前判断无解可能 | 组合问题剩余元素不足 |
| 条件剪枝 | 利用约束条件跳过无效分支 | N皇后冲突检测 |
| 记忆化剪枝 | 记录已搜索状态避免重复 | 带记忆化的递归搜索 |
4.2 剪枝实战优化
def combination_sum2(candidates, target):
"""LeetCode 40: 组合总和II - 每个数只能用一次"""
candidates.sort() # 排序是剪枝的前提
result = []
def backtrack(start, path, remaining):
if remaining == 0:
result.append(path[:])
return
for i in range(start, len(candidates)):
# 剪枝1:当前元素已超过剩余目标,后面更大元素也超
if candidates[i] > remaining:
break
# 剪枝2:同层跳过重复元素
if i > start and candidates[i] == candidates[i - 1]:
continue
path.append(candidates[i])
backtrack(i + 1, path, remaining - candidates[i])
path.pop()
backtrack(0, [], target)
return result
# 示例
print(combination_sum2([10, 1, 2, 7, 6, 1, 5], 8))
# [[1,1,6], [1,2,5], [1,7], [2,6]]
5. 常见问题与解决方案
5.1 递归深度过大
问题:Python默认递归深度限制为1000,深度递归会触发 RecursionError
import sys
sys.setrecursionlimit(10000) # 增大递归深度限制
# 更好的方案:转换为迭代或使用显式栈
def dfs_iterative(root):
stack = [(root, False)]
while stack:
node, visited = stack.pop()
if node is None:
continue
if visited:
# 处理节点
process(node)
else:
stack.append((node.right, False))
stack.append((node.left, False))
stack.append((node, True))
5.2 回溯结果重复
问题:结果中出现重复解
解决方案:
- 排序 + 跳过相邻重复元素
- 使用
used数组标记已使用元素 - 确保同一层不选相同元素,不同层可以选
5.3 路径拷贝遗漏
问题:结果中所有解都相同
# 错误:直接引用path,后续修改会影响已保存的结果
result.append(path)
# 正确:拷贝当前路径
result.append(path[:]) # 列表拷贝
result.append(path.copy()) # 列表拷贝
result.append(list(path)) # 列表拷贝
6. 总结与最佳实践
6.1 递归设计原则
- 明确基线条件:确保递归能够终止
- 信任递归:假设子问题已正确解决,只关注当前层逻辑
- 注意栈空间:递归深度受调用栈限制
- 避免重复计算:使用记忆化或改为自底向上DP
6.2 回溯算法框架选择
- 子集/组合问题:用
start参数控制选择范围,避免重复 - 排列问题:用
used数组标记已使用元素 - 含重复元素:先排序,同层跳过重复
6.3 优化方向
- 剪枝:排序 + 提前终止是最常用的优化
- 状态压缩:用位运算代替
used数组 - 记忆化:对重复子问题缓存结果
- 迭代化:深度过大时转为迭代实现