前置知识: 算法与数据结构

递归与回溯

4 minIntermediate2026/6/13

递归思想、回溯算法框架、经典回溯问题(子集、排列、组合、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 递归的执行过程

递归的执行分为两个阶段:

  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 递归设计原则

  1. 明确基线条件:确保递归能够终止
  2. 信任递归:假设子问题已正确解决,只关注当前层逻辑
  3. 注意栈空间:递归深度受调用栈限制
  4. 避免重复计算:使用记忆化或改为自底向上DP

6.2 回溯算法框架选择

  • 子集/组合问题:用 start 参数控制选择范围,避免重复
  • 排列问题:用 used 数组标记已使用元素
  • 含重复元素:先排序,同层跳过重复

6.3 优化方向

  1. 剪枝:排序 + 提前终止是最常用的优化
  2. 状态压缩:用位运算代替 used 数组
  3. 记忆化:对重复子问题缓存结果
  4. 迭代化:深度过大时转为迭代实现