CS1602计算导论
第 7 讲Part 3 问题求解与 AI 协作AI Level 0

递归

Recursion

一个函数可以调用它自己。这个听起来像绕口令的想法,是解决一大类问题最直接的办法——只要你能把大问题说成"小一号的同一个问题"。

本讲结束后你应当能

  • 分清递归函数的调用顺序和执行顺序
  • 说出递归的两个必备要素,并解释缺了 base case 会怎样
  • 用调用栈解释递归为什么会耗尽内存
  • 用递归实现分治:二分查找、归并排序、生成所有排列
  • 说清楚汉诺塔式递归和求值式递归的区别
本页目录

函数可以调用自己

你已经见过函数调用函数:

def square(x: int) -> int:
    return x * x

def sum_of_squares(a: int, b: int) -> int:
    return square(a) + square(b)      # 调用了另一个函数

print(sum_of_squares(3, 4))

现在把它推到极端——让函数调用它自己:

def factorial(n: int) -> int:
    if n <= 1:
        return 1                       # 到底了
    return n * factorial(n - 1)        # 调用自己,但参数小了一号

print(factorial(5))

这就是递归(recursion)。它能成立,靠的是一个观察:

5 的阶乘 = 5 × 4 的阶乘 4 的阶乘 = 4 × 3 的阶乘 …… 1 的阶乘 = 1(不用再往下问了)

把”求 5 的阶乘”这个问题,变成了”求 4 的阶乘”这个更小的同类问题。 一直缩小,直到小得可以直接回答。

数学里你早就见过这种定义方式。斐波那契数列的定义 F(n) = F(n-1) + F(n-2) 就是递归的——用自己定义自己。

调用顺序与执行顺序

这是理解递归最关键的一步,也是最容易糊涂的地方。调用是一路向下的,计算是一路向上的。

给上面的阶乘加上追踪:

def factorial(n: int, depth: int = 0) -> int:
    pad = "  " * depth
    print(f"{pad}调用 factorial({n})")

    if n <= 1:
        print(f"{pad}到底了,返回 1")
        return 1

    result = n * factorial(n - 1, depth + 1)

    print(f"{pad}factorial({n}) 算出 {result}")
    return result

factorial(5)

看输出的形状:先一路缩进下去(调用),到底之后再一路退回来(计算)。

关键在于 n * factorial(n - 1) 这一行:要算出乘法,必须先知道 factorial(n-1) 是多少。所以它必须等着。等到最深处返回 1,上一层才能算 2 * 1,再上一层才能算 3 * 2……

调用栈

每次函数调用,Python 都要记住”我是从哪儿来的、局部变量是什么”。这些信息存在一个叫调用栈(call stack)的地方。

栈的行为像一摞书:后放上去的先拿下来。factorial(5) 调用 factorial(4) 时,factorial(5) 还没结束,它被压在下面等着;factorial(4) 完成后弹出,factorial(5) 才继续。

调用 factorial(5)   →  栈: [f(5)]
调用 factorial(4)   →  栈: [f(5), f(4)]
调用 factorial(3)   →  栈: [f(5), f(4), f(3)]
调用 factorial(2)   →  栈: [f(5), f(4), f(3), f(2)]
调用 factorial(1)   →  栈: [f(5), f(4), f(3), f(2), f(1)]
f(1) 返回 1         →  栈: [f(5), f(4), f(3), f(2)]
f(2) 返回 2         →  栈: [f(5), f(4), f(3)]
...

栈会满

栈的空间有限。递归太深就会撑爆:

import sys
print("Python 默认的递归深度上限:", sys.getrecursionlimit())
def countdown(n: int) -> int:
    if n == 0:
        return 0
    return countdown(n - 1)

print(countdown(5000))

RecursionError 就是栈满了。这不是你的算法错了——是递归的深度超过了 Python 允许的层数。

递归的两个要素

任何递归函数都必须有这两样,缺一不可:

一、Base case(基准情形)——小到可以直接回答的情况,不再调用自己。

二、递推关系——把当前问题变成规模更小的同类问题。

def factorial(n: int) -> int:
    if n <= 1:          # ① base case
        return 1
    return n * factorial(n - 1)     # ② 递推,且 n-1 < n

第二点里”规模更小”很关键。如果每次调用参数没有变小,就永远到不了 base case:

def never_ends(n: int) -> int:
    if n == 0:
        return 0
    return never_ends(n)      # n 没变小
print(never_ends(3))

几个经典例子

最大公约数

欧几里得算法:gcd(a, b) = gcd(b, a % b),直到 b 为 0。

def gcd(a: int, b: int) -> int:
    if b == 0:
        return a
    return gcd(b, a % b)

print(gcd(48, 18))
print(gcd(1071, 462))

这是递归最漂亮的例子之一——三行代码,两千三百年历史。

斐波那契数列

def fib(n: int) -> int:
    if n <= 1:
        return n
    return fib(n - 1) + fib(n - 2)

print([fib(i) for i in range(10)])

定义直接翻译成代码,简洁得不像话。但它有个严重问题:

calls = 0

def fib(n: int) -> int:
    global calls
    calls += 1
    if n <= 1:
        return n
    return fib(n - 1) + fib(n - 2)

for n in [10, 20, 25]:
    calls = 0
    fib(n)
    print(f"fib({n}) 需要 {calls} 次函数调用")

fib(25) 要调用二十多万次。因为 fib(5) 被算了很多遍——fib(7) 要算它,fib(6) 也要算它,而这两个又各自被上层反复要求。

同一个子问题被重复计算了无数次。 解决办法是把算过的记下来,这叫记忆化,L11 会讲。现在你只要知道:递归写起来简洁,但可能藏着巨大的浪费。

快速幂

算 a的n次方,直接乘要乘 n 次。用递归可以只乘 log n 次:

def power(a: int, n: int) -> int:
    if n == 0:
        return 1
    half = power(a, n // 2)          # 只算一半
    if n % 2 == 0:
        return half * half
    return half * half * a

print(power(2, 10))
print(power(3, 5))
print(power(2, 100))

思路是:a^10 = (a^5)^2,而 a^5 = (a^2)^2 × a。每次把指数砍一半,所以只要 log n 步。

注意 half = power(a, n // 2) 这一行——先算出来存着,再用两次。如果写成 power(a, n//2) * power(a, n//2),就算了两遍,优势全没了。

分治

上面的快速幂用的是分治(divide and conquer)思想:把问题分成小块,分别解决,再合起来。递归是表达分治最自然的方式。

嵌套列表求和

一个列表里可能还套着列表,深度不定:

def deep_sum(items: list) -> int:
    total = 0
    for x in items:
        if isinstance(x, list):
            total += deep_sum(x)     # 是列表就递归进去
        else:
            total += x
    return total

print(deep_sum([1, [2, 3], [4, [5, [6, 7]]]]))

用循环写这个会很痛苦——你不知道要嵌几层。递归天然处理”深度不确定”的结构。

二分查找

在已排序的列表里找一个数。每次和中间的比一下,就能排除掉一半:

def binary_search(nums: list[int], target: int, lo: int = 0, hi: int | None = None) -> int:
    if hi is None:
        hi = len(nums) - 1
    if lo > hi:
        return -1                      # 找不到

    mid = (lo + hi) // 2
    if nums[mid] == target:
        return mid
    if nums[mid] < target:
        return binary_search(nums, target, mid + 1, hi)    # 找右半
    return binary_search(nums, target, lo, mid - 1)        # 找左半

data = [1, 3, 5, 7, 9, 11, 13, 15]
for t in [7, 1, 15, 8]:
    print(f"找 {t} → 下标 {binary_search(data, t)}")

归并排序

分治排序:把列表劈成两半,分别排好,再合并两个有序列表。

def merge(a: list[int], b: list[int]) -> list[int]:
    """合并两个已排序的列表"""
    result: list[int] = []
    i = j = 0
    while i < len(a) and j < len(b):
        if a[i] <= b[j]:
            result.append(a[i]); i += 1
        else:
            result.append(b[j]); j += 1
    result.extend(a[i:])
    result.extend(b[j:])
    return result

def merge_sort(nums: list[int]) -> list[int]:
    if len(nums) <= 1:                 # base case:0 个或 1 个自然有序
        return nums
    mid = len(nums) // 2
    left = merge_sort(nums[:mid])      # 排左半
    right = merge_sort(nums[mid:])     # 排右半
    return merge(left, right)          # 合起来

print(merge_sort([7, 3, 4, 5, 1, 2, 3]))

注意 merge_sort 的结构完全是递归三要素:base case 是”长度不超过 1”,递推是”两个半长的列表”,规模确实在缩小。

汉诺塔:一次思维的跳跃

前面所有的递归,返回的都是一个值——阶乘返回数字,二分查找返回下标,归并排序返回列表。

汉诺塔不一样。它返回的是”一件事被做完了”。

规则:三根柱子,n 个大小不同的盘子按大在下、小在上叠在第一根柱子上。每次只能移动一个盘子,且任何时候大盘不能压在小盘上。要把所有盘子移到第三根柱子。

直接想怎么移会绕晕。但换个角度:

想把 n 个盘子从 A 移到 C,只需要:

  1. 先把上面 n-1 个盘子从 A 移到 B(借道 C)
  2. 把最大的那个从 A 移到 C
  3. 再把 n-1 个盘子从 B 移到 C(借道 A)

第 1 步和第 3 步又是同一个问题,只是规模小了一号。

def hanoi(n: int, source: str, target: str, spare: str) -> None:
    """把 n 个盘子从 source 移到 target,spare 作为中转"""
    if n == 1:
        print(f"  盘子 1: {source} → {target}")
        return
    hanoi(n - 1, source, spare, target)          # ① 上面 n-1 个挪到中转柱
    print(f"  盘子 {n}: {source} → {target}")     # ② 最大的直接过去
    hanoi(n - 1, spare, target, source)          # ③ 再从中转柱挪过来

print("3 个盘子:")
hanoi(3, "A", "C", "B")

数一下,3 个盘子用了 7 步。你在 Lab 4 用循环算过这个数:2ⁿ - 1。

def hanoi(n: int, source: str, target: str, spare: str) -> None:
    if n == 1:
        return
    hanoi(n - 1, source, spare, target)
    hanoi(n - 1, spare, target, source)

def count_moves(n: int) -> int:
    moves = 0
    def go(k: int) -> None:
        nonlocal moves
        moves += 1
        if k > 1:
            go(k - 1); go(k - 1)
    go(n)
    return moves

for n in range(1, 8):
    print(f"{n} 个盘子: {count_moves(n)} 步 (2^{n} - 1 = {2**n - 1})")

生成所有可能

全排列

生成 1..n 的所有排列。思路:每个位置轮流放每个还没用过的数。

def permutations(items: list[int]) -> list[list[int]]:
    if len(items) <= 1:
        return [items]                      # base case

    result: list[list[int]] = []
    for i in range(len(items)):
        rest = items[:i] + items[i + 1:]    # 去掉第 i 个
        for p in permutations(rest):        # 剩下的全排列
            result.append([items[i]] + p)   # 把第 i 个放到最前面
    return result

for p in permutations([1, 2, 3]):
    print(p)

标准库里有现成的,可以拿来对照验证自己写得对不对:

from itertools import permutations as std_perm

mine = [[1,2,3],[1,3,2],[2,1,3],[2,3,1],[3,1,2],[3,2,1]]
std = [list(p) for p in std_perm([1, 2, 3])]
print(sorted(mine) == sorted(std))

八皇后

在 8×8 棋盘上放 8 个皇后,互相不能攻击(不同行、不同列、不同对角线)。

因为每行必须恰好放一个,可以用一个列表表示解:queens[i] = j 表示第 i 行的皇后在第 j 列。

def is_safe(queens: list[int], col: int) -> bool:
    """已放好的皇后中,是否有和新皇后冲突的"""
    row = len(queens)
    for r, c in enumerate(queens):
        if c == col:                        # 同列
            return False
        if abs(row - r) == abs(col - c):    # 同对角线
            return False
    return True

def solve(n: int, queens: list[int] | None = None) -> int:
    if queens is None:
        queens = []
    if len(queens) == n:
        return 1                            # 放满了,找到一个解
    count = 0
    for col in range(n):
        if is_safe(queens, col):
            count += solve(n, queens + [col])   # 放下这个,继续往下
    return count

for n in range(4, 9):
    print(f"{n} 皇后: {solve(n)} 个解")

这个套路叫回溯(backtracking):往下试,试不通就退回来换一个。它本质上就是递归——solve 负责”在已有的部分解基础上,数出所有完整解”。

一个能跑的分形

分形是”自己包含自己的缩小版”的图形,天然适合递归。谢尔宾斯基三角形可以用文本画出来:

def sierpinski(n: int) -> list[str]:
    """返回 2^n 行的谢尔宾斯基三角形"""
    if n == 0:
        return ["*"]

    smaller = sierpinski(n - 1)
    size = len(smaller)
    pad = " " * size
    # 上半:一个小三角居中;下半:两个小三角并排
    top = [pad + line + pad for line in smaller]
    bottom = [line + " " + line for line in smaller]
    return top + bottom

for line in sierpinski(4):
    print(line)

改一下 sierpinski(4) 里的数字看看——每加 1,图形的边长翻倍,而代码一个字都不用改。

递归还是循环

任何递归都能改写成循环,反之亦然。 那什么时候用哪个?

递归循环
代码长度通常更短通常更长
可读性问题本身是递归的时候更清楚问题是线性的时候更清楚
开销每层都要压栈,有额外开销没有
深度限制有(默认 1000 层)没有
# 阶乘:循环版更直白,也更省
def factorial_loop(n: int) -> int:
    result = 1
    for i in range(2, n + 1):
        result *= i
    return result

def factorial_rec(n: int) -> int:
    return 1 if n <= 1 else n * factorial_rec(n - 1)

print(factorial_loop(10), factorial_rec(10))

阶乘、求和这类线性问题,循环更合适——你能清楚说出”从 i 到 i+1 怎么变”。

但汉诺塔、八皇后、嵌套结构遍历这类问题,用循环写会极其痛苦,因为你得自己维护一个栈来模拟递归。

小结

  • 递归 = 把问题变成小一号的同类问题,直到小得可以直接回答
  • 调用一路向下,计算一路向上
  • 两个必备要素:base case 和规模递减的递推。缺一个就是 RecursionError
  • 每层调用都占调用栈空间,Python 默认限制约 1000 层
  • 写递归时做信仰之跃:假设下一层会正确完成它的活
  • 朴素的 fib 会重复计算同一个子问题成千上万次——L11 讲怎么修
  • 分治:二分查找、归并排序,都是”分成小块、分别解决、再合起来”
  • 汉诺塔式递归返回的不是值,而是”一件事做完了”——这是思维的门槛
  • 回溯 = 往下试,不通就退回来换一个
  • 结构本身递归的问题用递归,线性的问题用循环

练习

  1. 写 count_digits(n: int) -> int,用递归数一个非负整数有几位。想清楚 n = 0 时的 base case。
  2. 写 reverse_string(s: str) -> str,用递归反转字符串。不许用切片 [::-1]。
  3. 写 sum_digits(n: int) -> int,用递归求各位数字之和。和你在 Lab 4 写的循环版本对比一下。
  4. 给 fib 加一个计数器,画出 n 从 1 到 25 时调用次数的增长。你觉得它大概是什么增长速度?
  5. 写 flatten(items: list) -> list[int],把任意深度的嵌套列表压平成一维。 例如 [1, [2, [3, [4]]]] → [1, 2, 3, 4]。
  6. 修改本讲的 hanoi(),让它返回移动步骤的列表而不是打印。 想一想:这样改之后,它还是”完成一件事”式的递归吗?
  7. 用递归实现 is_palindrome(s: str) -> bool,判断字符串是否回文(正读反读一样)。

本讲的配套上机题在 Lab 7,个人项目 A 也在本周启动。