CS1602计算导论
第 11 讲Part 5 健壮性与真实世界AI Level 1

复杂度、异常与测试

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 变大时,它们的增长速度完全不同。

nn + 2n²
1012100
10010210,000
1,0001,0021,000,000
1,000,0001,000,0021,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 list1,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)。

怎么改进复杂度

按有效性排序:

  1. 换算法——O(n²) 变 O(n log n),这是数量级的改进
  2. 换数据结构——列表查找换成集合,O(n) 变 O(1)
  3. 避免重复计算——记忆化
  4. 用内置函数——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 主动抛出;自定义异常继承 Exception
  • with 保证清理,等价于写好的 try...finally
  • Python 偏好 EAFP(先做再道歉)而不是 LBYL

断言与测试

  • 异常处理外部的、可预期的问题;断言检查内部的、不该发生的情况
  • 不要捕获 AssertionError,断言失败就该崩
  • 断言在 -O 模式下会被跳过,所以不能用来校验用户输入
  • pytest:test_ 开头的文件和函数,用 assert 表达期望,pytest.raises 测异常
  • 实验平台就是调用你的函数比对返回值——所以必须 return
  • 每个函数至少测:正常、空、单个、边界、异常

练习

  1. 判断下面每个函数的时间复杂度,并说明理由:
    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
  2. 把 L7 里你写的 count_ways(硬币组合)加上记忆化,测量加速比。
  3. 写一个 safe_int(text: str, default: int = 0) -> int,转换失败时返回默认值而不是崩溃。
  4. 下面的代码有什么问题?改对它:
    def process(filename):
        try:
            f = open(filename)
            data = f.read()
            return int(data)
        except:
            return 0
  5. 给 Lab 9 你写的 Stack 类补一套 pytest 测试,至少覆盖:空栈、单元素、多元素、空栈 pop 的行为。
  6. 用 assert 给 L7 的 binary_search 加上前置条件和后置条件(提示:后置条件可以检查”返回的下标处确实是 target,或者返回 -1 且 target 确实不在里面”)。
  7. 实测一下:用 time 测量 list.insert(0, x) 和 list.append(x) 各执行 10000 次的耗时,验证 O(n) 和 O(1) 的差别。

本讲的配套上机题在 Lab 11。