可迭代对象与生成器
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、range | iter(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))
看清楚这个执行顺序:
- 调用
demo()不执行任何代码,只是造出一个生成器对象 - 第一次
next()才开始执行,跑到第一个yield就暂停,把值交出来 - 下一次
next()从暂停的地方继续,跑到下一个yield - 函数结束时抛
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)
练习
- 写一个生成器
evens(n),产出前 n 个偶数。再写一个squares(gen),接收一个生成器并产出每个值的平方。组合使用它们。 - 用生成器实现”读一个文件,产出所有非空行、去掉首尾空白”。为什么这里用生成器比返回列表好?
- 下面的代码为什么第二次打印是空的?改对它:
data = (x for x in range(5)) print(sum(data)) print(max(data)) - 写一个无限生成器
primes()产出所有质数,用itertools.islice取前 20 个。 - 用
collections.Counter重写 Lab 6 的词频统计,对比代码长度。 - 用
deque实现一个”最近 N 条记录”的缓冲区:加入新记录时,超过 N 条就自动丢掉最老的。(提示:deque(maxlen=N)) - 解释下面两段的内存差别,并用
sys.getsizeof验证:a = [x ** 2 for x in range(1_000_000)] b = (x ** 2 for x in range(1_000_000))
本讲的配套上机题在 Lab 13。