递归
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,只需要:
- 先把上面 n-1 个盘子从 A 移到 B(借道 C)
- 把最大的那个从 A 移到 C
- 再把 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 讲怎么修 - 分治:二分查找、归并排序,都是”分成小块、分别解决、再合起来”
- 汉诺塔式递归返回的不是值,而是”一件事做完了”——这是思维的门槛
- 回溯 = 往下试,不通就退回来换一个
- 结构本身递归的问题用递归,线性的问题用循环
练习
- 写
count_digits(n: int) -> int,用递归数一个非负整数有几位。想清楚n = 0时的 base case。 - 写
reverse_string(s: str) -> str,用递归反转字符串。不许用切片[::-1]。 - 写
sum_digits(n: int) -> int,用递归求各位数字之和。和你在 Lab 4 写的循环版本对比一下。 - 给
fib加一个计数器,画出n从 1 到 25 时调用次数的增长。你觉得它大概是什么增长速度? - 写
flatten(items: list) -> list[int],把任意深度的嵌套列表压平成一维。 例如[1, [2, [3, [4]]]]→[1, 2, 3, 4]。 - 修改本讲的
hanoi(),让它返回移动步骤的列表而不是打印。 想一想:这样改之后,它还是”完成一件事”式的递归吗? - 用递归实现
is_palindrome(s: str) -> bool,判断字符串是否回文(正读反读一样)。
本讲的配套上机题在 Lab 7,个人项目 A 也在本周启动。