CS1602计算导论
第 13 讲Part 6 Pythonic 与现代实践AI Level 1

可迭代对象与生成器

Pythonic I: Iterable and Generator

你写了十二讲 for 循环,但一直不知道它背后发生了什么。这一讲揭开它——顺便让你能处理比内存还大的数据。

本讲结束后你应当能

  • 说清楚 for 循环背后的迭代协议
  • 分辨可迭代对象和迭代器
  • 用 yield 写生成器,并说出它省下了什么
  • 用生成器表达式和推导式写出简洁的数据处理
  • 用 itertools 和 collections 解决常见问题
本页目录

for 循环背后

你从 L4 开始就在写 for x in something。它能作用于列表、字符串、字典、集合、range、甚至文件对象。

这些东西没有共同的父类。Python 是怎么统一处理它们的?

答案是一个约定,叫迭代协议(iteration protocol)。任何对象只要遵守这个约定,for 就能用在它身上——这正是 L10 讲的鸭子类型。

协议:两个函数

nums = [10, 20, 30]

it = iter(nums)          # 第一步:拿到一个迭代器
print(it)

print(next(it))          # 第二步:反复要下一个
print(next(it))
print(next(it))
nums = [10, 20, 30]
it = iter(nums)
for _ in range(3):
    next(it)
print(next(it))          # 取完了会怎样

取完之后 next() 抛出 StopIteration。

所以 for x in nums: 实际上等价于:

nums = [10, 20, 30]

# for x in nums: print(x)  展开后大致是这样
it = iter(nums)
while True:
    try:
        x = next(it)
    except StopIteration:
        break
    print(x)

for 只是这段代码的语法糖。它对每个对象做的事都一样:要一个迭代器,然后不停地 next,直到 StopIteration。

可迭代对象 vs 迭代器

两个词长得像,但不是一回事:

可迭代对象(iterable)迭代器(iterator)
定义实现了 __iter__实现了 __iter__ 和 __next__
能做什么能被 iter() 变成迭代器能被 next() 逐个取值
能否重复遍历能不能,取完就没了
例子list、str、dict、set、rangeiter(list) 的结果、生成器、文件对象
nums = [1, 2, 3]

# 列表可以反复遍历
print(list(nums), list(nums), list(nums))

# 迭代器只能遍历一次
it = iter(nums)
print(list(it))
print(list(it))          # 第二次是空的!
data = (x * 2 for x in range(5))
print("第一次:", sum(data))
print("第二次:", sum(data))

自己实现迭代协议

给 L9 的类加上迭代能力:

class Countdown:
    """从 n 倒数到 1。"""

    def __init__(self, n: int) -> None:
        self.n = n

    def __iter__(self) -> "CountdownIterator":
        return CountdownIterator(self.n)


class CountdownIterator:
    def __init__(self, n: int) -> None:
        self.current = n

    def __iter__(self) -> "CountdownIterator":
        return self                       # 迭代器返回自己

    def __next__(self) -> int:
        if self.current <= 0:
            raise StopIteration           # 取完了
        value = self.current
        self.current -= 1
        return value


c = Countdown(5)
for x in c:
    print(x, end=" ")
print()
for x in c:                               # 可以再遍历一次,因为 Countdown 是可迭代对象
    print(x, end=" ")
print()

写起来挺啰嗦的——两个类,一堆样板。生成器可以把这十几行压成三行。

生成器:yield

把 return 换成 yield,函数就变成了生成器函数:

def countdown(n: int):
    while n > 0:
        yield n                   # 交出一个值,然后暂停
        n -= 1

for x in countdown(5):
    print(x, end=" ")
print()

三行代码,和上面两个类做的事完全一样。

yield 和 return 的区别

return 结束函数,yield 暂停函数。

def demo():
    print("  开始")
    yield 1
    print("  第一次 yield 之后")
    yield 2
    print("  第二次 yield 之后")
    yield 3
    print("  结束")

g = demo()
print("生成器创建了,但函数体一行都没执行")
print("next →", next(g))
print("next →", next(g))
print("next →", next(g))

看清楚这个执行顺序:

  1. 调用 demo() 不执行任何代码,只是造出一个生成器对象
  2. 第一次 next() 才开始执行,跑到第一个 yield 就暂停,把值交出来
  3. 下一次 next() 从暂停的地方继续,跑到下一个 yield
  4. 函数结束时抛 StopIteration

函数的局部变量在暂停期间保持不变——这是生成器最神奇的地方。

def counter():
    count = 0
    while True:
        count += 1        # count 在两次 next 之间被记住了
        yield count

g = counter()
print(next(g), next(g), next(g), next(g))

生成器省下了什么

这才是重点。

import sys

# 列表:所有元素立刻造出来,全部装进内存
squares_list = [x ** 2 for x in range(100_000)]

# 生成器:一个都不造,用到才算
squares_gen = (x ** 2 for x in range(100_000))

print("列表占用:", sys.getsizeof(squares_list), "字节")
print("生成器占用:", sys.getsizeof(squares_gen), "字节")
print("两者的和相同:", sum(squares_list) == sum(squares_gen))

十万个元素的列表要几百 KB,生成器只要几十字节。

因为生成器不存数据,只存”怎么算下一个”的状态。这解释了一个你早就见过的东西:

import sys

r = range(10_000_000)
print("range 对象占用:", sys.getsizeof(r), "字节")
print("它代表一千万个数,却几乎不占内存")
print("因为它不是列表:", type(r))
print("需要哪个才算哪个:", r[5_000_000])

L4 说过”range 不是列表”,现在你知道为什么了。

处理比内存还大的数据

def read_numbers(lines: list[str]):
    """逐行产出数字,不把整个文件装进内存。"""
    for line in lines:
        line = line.strip()
        if line:
            yield int(line)

# 模拟一个"文件"
fake_file = ["1", "2", "", "3", "4", "5"]

total = 0
for n in read_numbers(fake_file):
    total += n
print("总和:", total)

真实场景里 lines 是一个文件对象。文件对象本身就是迭代器——这就是 L12 说”大文件用 for line in f 内存占用与文件大小无关”的原因。

无穷序列

生成器可以表示无限长的序列,因为它从不一次性生成全部:

def naturals():
    n = 1
    while True:
        yield n
        n += 1

def take(gen, k: int) -> list[int]:
    """从生成器里取前 k 个。"""
    result = []
    for x in gen:
        result.append(x)
        if len(result) == k:
            break
    return result

print(take(naturals(), 10))
def fibonacci():
    a, b = 0, 1
    while True:
        yield a
        a, b = b, a + b

fib = fibonacci()
print([next(fib) for _ in range(15)])

列表做不到这个——[n for n in naturals()] 会一直跑下去直到内存耗尽。

生成器表达式

和列表推导长得几乎一样,只是把方括号换成圆括号:

nums = range(1, 11)

as_list = [x ** 2 for x in nums]          # 列表推导:立刻造出全部
as_gen = (x ** 2 for x in nums)           # 生成器表达式:用到才算

print(as_list)
print(as_gen)
print(list(as_gen))

作为函数的唯一参数时,外面的括号可以省略:

nums = range(1, 1001)

print(sum(x ** 2 for x in nums))                    # 省略了括号
print(max(len(w) for w in ["a", "abc", "ab"]))
print(any(x > 900 for x in nums))
print(all(x > 0 for x in nums))

any 和 all 配生成器还有个额外好处——短路:

def checked(x: int) -> bool:
    print(f"  检查 {x}")
    return x > 3

print("any:", any(checked(x) for x in [1, 2, 3, 4, 5, 6]))
print("→ 找到第一个满足的就停了,后面的根本没检查")

各种推导式

words = ["apple", "banana", "cherry", "date"]

print([w.upper() for w in words])                     # 列表推导
print({w[0] for w in words})                          # 集合推导
print({w: len(w) for w in words})                     # 字典推导
print(tuple(len(w) for w in words))                   # 元组要用 tuple() 包

常见陷阱

1. 生成器只能用一次

gen = (x for x in range(5))
print(list(gen))
print(list(gen))          # 空的

需要多次用就转成列表:data = list(gen)。

2. 生成器是惰性的,异常会推迟

def risky():
    yield 1
    yield 2
    raise ValueError("出错了")
    yield 3

g = risky()
print("创建生成器时没有报错")
print(next(g))
print(next(g))
print(next(g))            # 现在才报

函数体里的错误,要等到那一行真的被执行时才暴露。 这让调试变得稍微麻烦一点。

3. 循环变量的延迟绑定

gens = [(x * i for x in range(3)) for i in range(3)]
print([list(g) for g in gens])

这个能正常工作,但换成 lambda 就不行了——那个坑在 L14 讲。

itertools:现成的迭代工具

标准库里有一整个模块,专门做迭代器的组合:

import itertools

print(list(itertools.count(10, 2))[:0] or "count 是无限的,要配 islice")
print(list(itertools.islice(itertools.count(10, 2), 5)))     # 从 10 开始,步长 2,取 5 个
print(list(itertools.islice(itertools.cycle("ABC"), 7)))     # 循环重复
print(list(itertools.repeat("x", 3)))                        # 重复 3 次
import itertools

print(list(itertools.chain([1, 2], [3, 4], [5])))            # 串联多个序列
print(list(itertools.accumulate([1, 2, 3, 4, 5])))           # 累加
print(list(itertools.combinations("ABC", 2)))                # 组合
print(list(itertools.permutations("ABC", 2)))                # 排列
print(list(itertools.product([1, 2], "AB")))                 # 笛卡尔积

groupby 常用来分组,但要求输入已排序:

import itertools

data = [("水果", "苹果"), ("水果", "香蕉"), ("蔬菜", "白菜"), ("蔬菜", "萝卜")]

for key, group in itertools.groupby(data, key=lambda x: x[0]):
    print(key, "→", [item[1] for item in group])
import itertools

data = ["苹果", "白菜", "香蕉"]
key = lambda w: "水果" if w in ("苹果", "香蕉") else "蔬菜"

print("没排序:", [(k, list(g)) for k, g in itertools.groupby(data, key=key)])
print("排序后:", [(k, list(g)) for k, g in itertools.groupby(sorted(data, key=key), key=key)])

collections:更好用的容器

from collections import Counter, defaultdict, deque

# Counter:计数,比手写 dict.get(k, 0) + 1 方便
words = "the quick brown fox the lazy dog the end".split()
c = Counter(words)
print(c)
print(c.most_common(2))
print(c["the"], c["不存在的词"])      # 不存在返回 0,不报 KeyError
from collections import defaultdict

# defaultdict:自动创建默认值,省掉 setdefault
groups = defaultdict(list)
for word in ["apple", "avocado", "banana", "blueberry"]:
    groups[word[0]].append(word)      # 不用先判断键在不在
print(dict(groups))

这正好解决了 L6 练习 6 那个 KeyError。

from collections import deque

# deque:两端都能高效增删(列表的 insert(0,x) 是 O(n),deque 是 O(1))
d = deque([1, 2, 3])
d.appendleft(0)
d.append(4)
print(d)
print(d.popleft(), d.pop())
print(d)
import time
from collections import deque

n = 50_000

lst = []
t = time.time()
for i in range(n):
    lst.insert(0, i)
list_t = time.time() - t

dq = deque()
t = time.time()
for i in range(n):
    dq.appendleft(i)
deque_t = time.time() - t

print(f"list.insert(0, x) {n} 次: {list_t:.4f} 秒")
print(f"deque.appendleft  {n} 次: {deque_t:.4f} 秒")
print(f"deque 快约 {list_t / deque_t:.0f} 倍")

这是 L11 说的”换数据结构”的一个实例:O(n) 换成 O(1)。

解包再看一眼

L6 讲过元组解包。* 和 ** 在函数调用里也能用:

def add(a: int, b: int, c: int) -> int:
    return a + b + c

nums = [1, 2, 3]
print(add(*nums))                # 把列表拆成三个参数

opts = {"a": 1, "b": 2, "c": 3}
print(add(**opts))               # 把字典拆成关键字参数
a, *rest = [1, 2, 3, 4]
print(a, rest)

first, *middle, last = [1, 2, 3, 4, 5]
print(first, middle, last)

print([*"abc", *[1, 2]])         # 解包进列表字面量
print({**{"a": 1}, **{"b": 2}})  # 解包进字典(L6 讲过的合并)

小结

  • for 是语法糖:iter() 拿迭代器,反复 next(),遇 StopIteration 停
  • 可迭代对象能反复遍历;迭代器只能遍历一次
  • yield 暂停函数,return 结束函数。局部变量在暂停期间保持
  • 生成器不存数据,只存怎么算下一个——内存占用与数据量无关
  • range 和文件对象都是这个道理
  • 生成器能表示无限序列
  • 生成器表达式 (...) 对上列表推导 [...]:只遍历一次就用生成器
  • any/all 配生成器会短路
  • 没有元组推导,(x for x in ...) 是生成器
  • itertools:islice、chain、accumulate、combinations、groupby(要先排序)
  • collections:Counter 计数、defaultdict 免判断、deque 两端 O(1)

练习

  1. 写一个生成器 evens(n),产出前 n 个偶数。再写一个 squares(gen),接收一个生成器并产出每个值的平方。组合使用它们。
  2. 用生成器实现”读一个文件,产出所有非空行、去掉首尾空白”。为什么这里用生成器比返回列表好?
  3. 下面的代码为什么第二次打印是空的?改对它:
    data = (x for x in range(5))
    print(sum(data))
    print(max(data))
  4. 写一个无限生成器 primes() 产出所有质数,用 itertools.islice 取前 20 个。
  5. 用 collections.Counter 重写 Lab 6 的词频统计,对比代码长度。
  6. 用 deque 实现一个”最近 N 条记录”的缓冲区:加入新记录时,超过 N 条就自动丢掉最老的。(提示:deque(maxlen=N))
  7. 解释下面两段的内存差别,并用 sys.getsizeof 验证:
    a = [x ** 2 for x in range(1_000_000)]
    b = (x ** 2 for x in range(1_000_000))

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