复杂度、异常与测试
Complexity, Exception and Testing
前面的程序都假设输入合法、数据干净、规模不大。这一讲处理真实世界——程序会用三种方式辜负你:太慢、会崩、悄悄算错。
本讲结束后你应当能
- 用大 O 记号描述算法的增长速度,并说出常见操作的复杂度
- 用记忆化把指数级的递归降到线性
- 用 try / except / else / finally 处理异常,并知道什么时候不该捕获
- 分清异常和 bug,正确使用 assert
- 用 pytest 写测试,并理解实验平台是怎么给你判分的
本页目录
程序辜负你的三种方式
到目前为止,你写的程序都活在一个理想世界里:输入总是合法的,数据总是干净的,规模总是很小。
真实世界不是这样。程序会以三种方式让你失望:
| 方式 | 症状 | 本讲的应对 |
|---|---|---|
| 太慢 | 小数据没问题,真实数据跑不完 | 复杂度分析 |
| 会崩 | 遇到意外输入直接终止 | 异常处理 |
| 悄悄算错 | 不报错,但结果是错的 | 断言与测试 |
三件事看起来无关,其实是同一个问题的三个面:你怎么知道你的程序真的能用?
一、太慢:复杂度
为什么不用秒表
想知道一个算法快不快,最直接的想法是跑一下计时:
import time
def sum_to(n: int) -> int:
total = 0
for i in range(n):
total += i
return total
for n in [10_000, 100_000, 1_000_000]:
t = time.time()
sum_to(n)
print(f"n = {n:>9,} 耗时 {time.time() - t:.4f} 秒")
这个数字有用,但它不能拿来比较算法,因为它取决于:机器快慢、当时 CPU 在忙什么、Python 版本、甚至你有没有开着别的程序。
我们需要一个与机器无关的度量。
数步数,而不是数秒
换个思路:数这个算法要执行多少次基本操作。基本操作指的是加减乘除、比较、赋值、访问下标这类”一步就能做完”的事。
def sum_to(n):
total = 0 # 1 步
for i in range(n): # 循环 n 次
total += i # 每次 1 步
return total # 1 步
总共大约 n + 2 步。
再看一个:
def has_duplicate(items):
for i in range(len(items)): # n 次
for j in range(len(items)): # 每次又 n 次
if items[i] == items[j]: # 1 步
...
这是 n × n = n² 步。
关键差别是:当 n 变大时,它们的增长速度完全不同。
| n | n + 2 | n² |
|---|---|---|
| 10 | 12 | 100 |
| 100 | 102 | 10,000 |
| 1,000 | 1,002 | 1,000,000 |
| 1,000,000 | 1,000,002 | 1,000,000,000,000 |
n 翻十倍,第一个也翻十倍,第二个翻一百倍。
大 O 记号
我们只关心 n 很大时的增长趋势,所以:
- 只保留增长最快的那一项:
n² + 3n + 5→ 看n² - 忽略常数系数:
5n和100n都记作n
写成 O(...),读作”大 O”:
n + 2→ O(n),线性n²→ O(n²),平方3→ O(1),常数(和 n 无关)
为什么可以忽略常数?因为增长的阶数一旦不同,常数再大也追不上:
# O(n) 但常数很大(每步 100 个操作)vs O(n²) 但常数很小
for n in [10, 100, 1000, 10000]:
linear = 100 * n
quadratic = n * n
winner = "O(n) 更少" if linear < quadratic else "O(n²) 更少"
print(f"n={n:>6} 100n = {linear:>12,} n² = {quadratic:>14,} {winner}")
n 超过 100 之后,O(n) 就再也不会输了。这就是为什么阶数比常数重要。
常见的复杂度
从快到慢:
| 记号 | 名称 | 例子 | n=1000 时大约几步 |
|---|---|---|---|
| O(1) | 常数 | 访问 a[i]、字典查找 | 1 |
| O(log n) | 对数 | 二分查找 | 10 |
| O(n) | 线性 | 遍历列表、x in list | 1,000 |
| O(n log n) | 线性对数 | 归并排序、sorted() | 10,000 |
| O(n²) | 平方 | 双重循环 | 1,000,000 |
| O(2ⁿ) | 指数 | 朴素递归斐波那契 | 天文数字 |
import math
print(f"{'n':>8} {'log n':>8} {'n':>10} {'n log n':>12} {'n²':>14}")
for n in [10, 100, 1000, 10000, 100000]:
print(f"{n:>8} {math.log2(n):>8.1f} {n:>10,} {n*math.log2(n):>12,.0f} {n*n:>14,}")
Python 各种操作的复杂度
这张表很实用,值得记住:
| 操作 | 复杂度 | 说明 |
|---|---|---|
a[i] | O(1) | 列表、元组、字符串按下标访问 |
a.append(x) | O(1) | 末尾追加 |
a.pop() | O(1) | 删末尾 |
a.insert(0, x) | O(n) | 后面全部要挪 |
a.pop(0) | O(n) | 同上 |
x in a(列表) | O(n) | 从头扫 |
x in s(集合/字典) | O(1) | 哈希 |
a.sort() / sorted(a) | O(n log n) | |
d[k] | O(1) | 字典查找 |
len(a) | O(1) | 长度是存好的,不用数 |
a + b(列表拼接) | O(n+m) | 造新列表 |
a[i:j](切片) | O(j−i) | 造新列表 |
s1 + s2(字符串) | O(n+m) | 字符串不可变,每次都造新的 |
这解释了 L6 讲容器选型时那句”列表查找慢、集合查找快”:
import time
n = 30_000
data_list = list(range(n))
data_set = set(data_list)
t = time.time()
for x in range(0, n, 100):
x in data_list
list_time = time.time() - t
t = time.time()
for x in range(0, n, 100):
x in data_set
set_time = time.time() - t
print(f"列表 in: {list_time:.4f} 秒")
print(f"集合 in: {set_time:.6f} 秒")
print(f"集合快约 {list_time / set_time:.0f} 倍")
import time
words = ["hello"] * 20_000
t = time.time()
s = ""
for w in words:
s += w
concat_time = time.time() - t
t = time.time()
s2 = "".join(words)
join_time = time.time() - t
print(f"+= 拼接: {concat_time:.4f} 秒")
print(f"join: {join_time:.6f} 秒")
print(f"结果一样吗: {s == s2}")
二分查找为什么是 O(log n)
L7 写过二分查找,当时说”复杂度留到 L11”。现在补上。
每比较一次,候选范围就减半:n → n/2 → n/4 → … → 1。
问:除以多少次 2 才能从 n 降到 1? 答案就是 log₂ n。
import math
for n in [100, 1_000, 1_000_000, 1_000_000_000]:
print(f"{n:>15,} 个元素,最多比较 {math.ceil(math.log2(n)):>2} 次")
十亿个元素,三十次比较。 这就是对数复杂度的威力,也是为什么排好序的数据这么值钱。
记忆化:把指数降到线性
L7 里我们发现朴素的斐波那契要调用二十多万次,并说”L11 讲怎么修”。
问题的根源是同一个子问题被算了无数遍。解决办法很直接:算过的记下来。
import time
def fib_naive(n: int) -> int:
if n <= 1:
return n
return fib_naive(n - 1) + fib_naive(n - 2)
def fib_memo(n: int, cache: dict[int, int] | None = None) -> int:
if cache is None:
cache = {}
if n in cache:
return cache[n] # 算过了,直接取
if n <= 1:
return n
result = fib_memo(n - 1, cache) + fib_memo(n - 2, cache)
cache[n] = result # 记下来
return result
t = time.time(); fib_naive(28); naive_t = time.time() - t
t = time.time(); fib_memo(28); memo_t = time.time() - t
print(f"朴素递归 fib(28): {naive_t:.4f} 秒")
print(f"记忆化 fib(28): {memo_t:.6f} 秒")
print(f"快约 {naive_t / memo_t:.0f} 倍")
print(f"记忆化能轻松算 fib(200) = {fib_memo(200)}")
复杂度从 O(2ⁿ) 降到 O(n)——每个 n 只算一次。
Python 标准库有个现成的装饰器(L14 会讲装饰器是什么):
from functools import lru_cache
@lru_cache(maxsize=None)
def fib(n: int) -> int:
if n <= 1:
return n
return fib(n - 1) + fib(n - 2)
print(fib(100))
print(fib.cache_info())
加一行 @lru_cache,指数变线性。
空间复杂度
同样的记号也用来描述内存占用。
# O(1) 空间:不管 n 多大,只用几个变量
def sum_loop(n: int) -> int:
total = 0
for i in range(n):
total += i
return total
# O(n) 空间:造了一个 n 长的列表
def sum_list(n: int) -> int:
return sum([i for i in range(n)])
print(sum_loop(1000), sum_list(1000))
L13 讲生成器时,你会看到怎么把 O(n) 空间降回 O(1)。
怎么改进复杂度
按有效性排序:
- 换算法——O(n²) 变 O(n log n),这是数量级的改进
- 换数据结构——列表查找换成集合,O(n) 变 O(1)
- 避免重复计算——记忆化
- 用内置函数——
sum()、sorted()、"".join()底层是 C 写的,比纯 Python 循环快一个数量级
import time
n = 2_000_000
t = time.time()
total1 = 0
for i in range(n):
total1 += i
loop_t = time.time() - t
t = time.time()
total2 = sum(range(n))
builtin_t = time.time() - t
print(f"手写循环: {loop_t:.4f} 秒")
print(f"内置 sum: {builtin_t:.4f} 秒")
print(f"快约 {loop_t / builtin_t:.1f} 倍(同为 O(n),差在常数)")
二、会崩:异常
语法对了,还是会出错
x = 10
y = 0
print(x / y)
这段代码语法完全正确,Python 也顺利编译了,但运行时出错了。这种运行时才发生的错误叫异常(exception)。
不处理的话,程序直接终止——后面的代码一行都不会执行。
try / except
def safe_divide(a: float, b: float) -> float | None:
try:
return a / b
except ZeroDivisionError:
print("除数不能为零")
return None
print(safe_divide(10, 2))
print(safe_divide(10, 0))
print("程序继续运行")
try 里的代码正常执行;一旦抛出异常,立刻跳到匹配的 except 分支。
捕获具体的异常
def parse_age(text: str) -> int | None:
try:
return int(text)
except ValueError:
print(f"「{text}」不是一个整数")
return None
for s in ["18", "abc", "3.14", ""]:
print(s, "→", parse_age(s))
这正好解决了 L3 留下的问题——那时候我说”输入非数字程序会崩,L11 教你怎么办”。
多个 except
def get_item(data: dict[str, list[int]], key: str, index: int) -> int | None:
try:
return data[key][index]
except KeyError:
print(f"没有键 {key}")
except IndexError:
print(f"下标 {index} 越界")
except TypeError as e:
print(f"类型不对: {e}")
return None
d = {"a": [1, 2, 3]}
print(get_item(d, "a", 1))
print(get_item(d, "b", 0))
print(get_item(d, "a", 10))
最多只有一个分支会执行,按从上到下第一个匹配的算。
else 和 finally
def divide(a: float, b: float) -> None:
try:
result = a / b
except ZeroDivisionError:
print(" 出错了")
else:
print(f" 成功,结果 {result}") # 没出异常时才执行
finally:
print(" 收尾工作") # 无论如何都执行
print("10 / 2:"); divide(10, 2)
print("10 / 0:"); divide(10, 0)
else:try里没有抛异常时执行。用来放”成功之后才该做的事”finally:不管怎样都执行。用来放清理工作——关文件、断开连接
finally 连 return 都拦不住:
def f() -> str:
try:
return "从 try 返回"
finally:
print(" finally 仍然执行了")
print(f())
raise:主动抛出
def set_age(age: int) -> int:
if age < 0:
raise ValueError(f"年龄不能为负: {age}")
if age > 150:
raise ValueError(f"年龄不合理: {age}")
return age
print(set_age(18))
print(set_age(-5))
什么时候该 raise 而不是返回 None?
当调用方不应该忽略这个问题的时候。返回 None 很容易被忽略掉,然后错误传播到很远的地方才爆发;抛异常则立刻停下,traceback 直接指向出问题的位置。
自定义异常
class InsufficientFunds(Exception):
"""余额不足。"""
def __init__(self, balance: float, amount: float) -> None:
super().__init__(f"余额 {balance},无法取出 {amount}")
self.balance = balance
self.amount = amount
def withdraw(balance: float, amount: float) -> float:
if amount > balance:
raise InsufficientFunds(balance, amount)
return balance - amount
try:
withdraw(100, 500)
except InsufficientFunds as e:
print("捕获到:", e)
print("还差:", e.amount - e.balance)
自定义异常继承自 Exception(L10 的继承派上用场了)。好处是调用方能精确地捕获这一种错误,而不是笼统地抓 ValueError。
with:自动清理
# 手动版本,容易忘记关
f = open("/tmp/demo.txt", "w")
f.write("hello")
f.close()
# with 版本,自动关
with open("/tmp/demo.txt", "w") as f:
f.write("hello")
# 到这里文件已经关好了,即使中途出异常也一样
with open("/tmp/demo.txt") as f:
print(f.read())
with 保证退出时执行清理,等价于一个写好的 try...finally。文件操作的完整内容在 L12。
该检查还是该捕获
两种风格:
d = {"a": 1}
# 风格一:先看再跳(Look Before You Leap)
if "b" in d:
print(d["b"])
else:
print("LBYL: 没有这个键")
# 风格二:先做再道歉(Easier to Ask Forgiveness than Permission)
try:
print(d["b"])
except KeyError:
print("EAFP: 没有这个键")
Python 社区偏好第二种(EAFP)。理由:
- 正常路径没有额外开销(异常只在真出错时才有代价)
- 避免”检查完到使用之间状态变了”的竞争问题
- 代码主干更清晰——正常流程一眼可见,异常处理在旁边
但简单的检查该做还是做。if not items: return 比包一层 try 清楚得多。
三、悄悄算错:断言与测试
assert:检查你以为一定成立的事
def average(nums: list[float]) -> float:
assert len(nums) > 0, "average 不接受空列表"
return sum(nums) / len(nums)
print(average([1, 2, 3]))
print(average([]))
assert 条件, 消息 的意思是:“我认为这里条件一定成立。如果不成立,说明我的程序有 bug,立刻停下。”
断言不是异常处理
这是最容易混淆的一点,务必分清:
| 异常(try/except) | 断言(assert) | |
|---|---|---|
| 用于 | 可预期的外部问题 | 不该发生的内部错误 |
| 例子 | 文件不存在、用户输入非法、网络断了 | 函数收到了不可能的参数 |
| 谁的错 | 环境或用户 | 程序员自己 |
| 该不该捕获 | 该 | 不该 |
| 生产环境 | 保留 | 可以关掉 |
print("__debug__ 的值:", __debug__)
print("这决定了 assert 会不会被执行。正常模式为 True,-O 模式为 False。")
用断言给代码写文档
def binary_search(sorted_nums: list[int], target: int) -> int:
assert sorted_nums == sorted(sorted_nums), "输入必须是已排序的"
lo, hi = 0, len(sorted_nums) - 1
while lo <= hi:
mid = (lo + hi) // 2
if sorted_nums[mid] == target:
return mid
if sorted_nums[mid] < target:
lo = mid + 1
else:
hi = mid - 1
return -1
print(binary_search([1, 3, 5, 7], 5))
print(binary_search([5, 3, 1], 3))
这条断言把”输入必须有序”这个前置条件写成了可执行的检查。比写在注释里强——注释会过期,断言不会。
pytest:把测试正规化
L8 说过:判断代码对不对的依据是测试,不是”读起来合理”。当时我们手写了一个检查循环。现在把它规范化。
安装:
python3 -m pip install pytest
写测试的规则很简单:
- 文件名以
test_开头 - 函数名以
test_开头 - 用
assert表达期望
假设你有 mymath.py:
# mymath.py
def add(a: int, b: int) -> int:
return a + b
def divide(a: float, b: float) -> float:
if b == 0:
raise ValueError("除数不能为零")
return a / b
对应的测试文件:
# test_mymath.py
import pytest
from mymath import add, divide
def test_add_positive():
assert add(2, 3) == 5
def test_add_negative():
assert add(-1, -1) == -2
def test_add_zero():
assert add(0, 5) == 5
def test_divide_normal():
assert divide(10, 2) == 5.0
def test_divide_by_zero():
with pytest.raises(ValueError): # 期望它抛异常
divide(10, 0)
然后:
python3 -m pytest # 跑所有测试
python3 -m pytest -v # 显示每个测试的名字
python3 -m pytest -k divide # 只跑名字含 divide 的
输出大致长这样:
test_mymath.py ..... [100%]
======== 5 passed in 0.02s ========
有失败的话,pytest 会告诉你哪个测试失败、期望什么、实际什么:
E assert 6 == 5
E + where 6 = add(2, 3)
在讲义里我们没法真的跑 pytest,但可以模拟它的核心思路:
def add(a: int, b: int) -> int:
return a + b
def run_tests() -> None:
cases = [
("正数相加", lambda: add(2, 3), 5),
("负数相加", lambda: add(-1, -1), -2),
("加零", lambda: add(0, 5), 5),
("故意写错", lambda: add(2, 3), 6), # 这个会失败
]
passed = 0
for name, func, expected in cases:
got = func()
if got == expected:
print(f" ✓ {name}")
passed += 1
else:
print(f" ✗ {name}: 期望 {expected},实际 {got}")
print(f"{passed}/{len(cases)} 通过")
run_tests()
该测什么
回顾 L8 的边界清单,一个函数至少该测:
- 正常情况:典型输入
- 空输入:空列表、空字符串
- 单个元素
- 边界值:0、负数、最大最小
- 异常情况:非法输入是否正确抛出
def clamp(x: float, lo: float, hi: float) -> float:
"""把 x 限制在 [lo, hi] 范围内。"""
if lo > hi:
raise ValueError("lo 不能大于 hi")
return max(lo, min(x, hi))
checks = [
("范围内", clamp(5, 0, 10), 5),
("小于下界", clamp(-3, 0, 10), 0),
("大于上界", clamp(99, 0, 10), 10),
("正好在边界", clamp(0, 0, 10), 0),
("上下界相等", clamp(7, 3, 3), 3),
]
for name, got, expected in checks:
print(f" {'✓' if got == expected else '✗'} {name}: {got}")
try:
clamp(5, 10, 0)
print(" ✗ 非法范围没有抛异常")
except ValueError:
print(" ✓ 非法范围正确抛出 ValueError")
小结
复杂度
- 用基本操作的步数衡量算法,与机器无关
- 大 O 只保留最高阶项,忽略常数——因为阶数不同时常数追不上
- 常见阶:O(1) < O(log n) < O(n) < O(n log n) < O(n²) < O(2ⁿ)
- 记住 Python 的操作复杂度:列表
in是 O(n),集合和字典是 O(1) - 循环里
+=拼字符串是 O(n²),要用join - 记忆化把重复子问题的递归从 O(2ⁿ) 降到 O(n),
@lru_cache一行搞定 - 内置函数快是常数小,救不了错的算法。先降阶,再优化常数
异常
- 语法正确的代码仍会在运行时出错,这叫异常
try/except 具体类型/else(没出错时)/finally(总是执行)- 永远不要写裸的
except: raise主动抛出;自定义异常继承Exceptionwith保证清理,等价于写好的try...finally- Python 偏好 EAFP(先做再道歉)而不是 LBYL
断言与测试
- 异常处理外部的、可预期的问题;断言检查内部的、不该发生的情况
- 不要捕获
AssertionError,断言失败就该崩 - 断言在
-O模式下会被跳过,所以不能用来校验用户输入 - pytest:
test_开头的文件和函数,用assert表达期望,pytest.raises测异常 - 实验平台就是调用你的函数比对返回值——所以必须
return - 每个函数至少测:正常、空、单个、边界、异常
练习
- 判断下面每个函数的时间复杂度,并说明理由:
def f1(n): return sum(range(n)) def f2(a): return [x for x in a if x in a] def f3(a): return sorted(set(a)) def f4(n): count = 0 i = n while i > 1: i = i // 2 count += 1 return count - 把 L7 里你写的
count_ways(硬币组合)加上记忆化,测量加速比。 - 写一个
safe_int(text: str, default: int = 0) -> int,转换失败时返回默认值而不是崩溃。 - 下面的代码有什么问题?改对它:
def process(filename): try: f = open(filename) data = f.read() return int(data) except: return 0 - 给 Lab 9 你写的
Stack类补一套 pytest 测试,至少覆盖:空栈、单元素、多元素、空栈 pop 的行为。 - 用
assert给 L7 的binary_search加上前置条件和后置条件(提示:后置条件可以检查”返回的下标处确实是 target,或者返回 -1 且 target 确实不在里面”)。 - 实测一下:用
time测量list.insert(0, x)和list.append(x)各执行 10000 次的耗时,验证 O(n) 和 O(1) 的差别。
本讲的配套上机题在 Lab 11。