LECTURE 10

迭代器与生成器

把「一串值」和「算到哪儿了」装进同一个对象里,于是序列可以无限长、可以只在需要时才被算出来——递归也可以一边算一边吐结果。

教材:Composing Programs §4.2 Implicit Sequences 对应作业:Lab 05 / HW 04

0. 本讲导读

到目前为止,你处理「一堆值」的方式只有一种:先把它们全部装进一个列表,再遍历这个列表。list_partitions(6, 4) 会先算出全部 9 个分法、拼成一个二维列表,然后你才能看第一个。写树的递归时也一样:count_paths 先把所有路径数完,才返回一个数。

这套做法有两个躲不开的代价。

第一,你必须付出全部计算,哪怕只想要第一个结果。list_partitions(60, 50) 有 966370 个分法,在我的机器上要跑将近 3 秒、并且把这 96 万个字符串同时留在内存里——而如果你只是想看看前 3 个长什么样,这 3 秒和这块内存全都白花了。

第二,有些序列根本装不进列表。「所有偶数」「斐波那契数列的每一项」「一个不断产生随机数的流」,它们没有末尾,list() 会一直跑到内存耗尽。

本讲给出的解决办法,是把「一串值」这件事拆成两层:

  • 可迭代对象(iterable)——「这里有一串值」。列表、元组、字符串、字典、range 都是。
  • 迭代器(iterator)——「我读到第几个了」。它是一个独立的对象,同时记着「值从哪来」和「当前位置」,你每次向它要一个值,它才去算/取一个值。

把位置信息独立成一个对象,序列就不必真的存在于内存里了:只要有人能回答「下一个是什么」,就够了。生成器(generator)正是让你用写函数的方式来回答这个问题——把 return 换成 yield,一个普通函数就变成了一台「按需产值的机器」,而且它能递归。本讲最后那个 yield_partitions 只有 6 行,却能在 0 秒内吐出 60 = 10 + 50 这第一个分法,而列表版要跑 3 秒。

与前后讲的关系:上一讲讲树(tree)的递归,本讲末尾的 yield_paths 就是把上一讲的 count_paths 从「数出来」改成「一条条吐出来」——递归的骨架一模一样,只是返回方式变了。下一讲讲异常(exception)II,而 StopIteration 正是本讲会反复见到的一个异常,两讲直接接得上。教材对应 §4.2 Implicit Sequences,作业是 Lab 05 与 HW 04。

核心结论
  • iterable 是「书」,iterator 是「书签」。 对 iterable 调 iter 得到一个 iterator;对 iterator 调 next 得到下一个值;没有下一个了就抛 StopIteration。
  • 迭代器是一次性的、有状态的。 已经 next 过的元素永远回不去;把一个迭代器传给别的函数,位置会跟着一起传过去。
  • 所有迭代器都是可迭代对象,反之不成立。 iter(iterable) 返回一个新的迭代器;iter(iterator) 返回它自己(同一个对象)。
  • map / filter / zip / reversed 返回的都是惰性(lazy)的迭代器:不调 next,里面的函数一次都不会被调用。要看全部内容就套一层 list / tuple / sorted。
  • 函数体里出现 yield,这个函数就变成生成器函数。 调用它不执行任何函数体代码,只返回一个生成器对象;每次 next 才把函数体从上次暂停的地方恢复,跑到下一个 yield 停住。
  • yield from <iterable> 把那个 iterable 里的元素逐个让出去。递归生成器必须用 yield from;写成 yield 会让出一个生成器对象本身,结果里就会出现 <generator object ...>。

1. 复习:for 循环里到底放了什么

本讲要拆开 for 循环的内部机制,所以先把你已经会用的两件事摆出来,等下好对照。

字典可以直接放进 for 循环

你已经知道用 .keys()、.values()、.items() 遍历字典。其实字典本身就能直接进 for,此时遍历的是键:

>>> roman_numerals = {'I': 1, 'V': 5, 'X': 10}
>>> for key in roman_numerals:
...     print(key)
...
I
V
X

也就是说 for key in d 和 for key in d.keys() 效果相同。这件小事在本讲的意义是:「能放进 for 循环」是一种资格,字典有、列表有、字符串有、range 有,而它们内部的存储方式完全不同(字典是哈希表,range 连元素都不存)。既然存储方式各不相同,for 却能一视同仁地遍历它们,说明 for 依赖的不是「存储方式」,而是某个统一的接口。这个接口就是本讲的主角。

注意

遍历字典的过程中不能增删键,否则 Python 会当场报错:

>>> for k in roman_numerals:
...     roman_numerals['L'] = 50
...
Traceback (most recent call last):
  ...
RuntimeError: dictionary changed size during iteration

报错信息说得很直白:迭代过程中字典的大小变了。为什么 Python 要专门检查这个?因为遍历字典时有一个「读到哪儿了」的位置,而增删键可能触发哈希表重排,位置就失去意义了。列表没有这层检查——第 3 节会看到,正因为没有检查,列表在迭代中被改会得到更诡异的结果。

树的遍历顺序:同样三行代码,位置一换结果全变

上一讲的树 ADT 继续沿用(本讲的随堂代码里原样带着它):

def tree(root_label, branches=[]):
    for branch in branches:
        assert is_tree(branch), 'branches must be trees'
    return [root_label] + list(branches)

def label(tree):
    return tree[0]

def branches(tree):
    return tree[1:]

def is_leaf(tree):
    return not branches(tree)

下面三个函数长得几乎一样,只差 print 放在哪儿。用同一棵树 t = tree(3, [tree(-1), tree(1, [tree(2, [tree(1)]), tree(3)]), tree(1, [tree(-1)])]) 跑一遍:

函数代码对 t 的输出它到底在做什么
print_treeA
print(label(tree))
for b in branches(tree):
    print_treeA(b)
3 -1 1 2 1 3 1 -1 先序:每个节点恰好打印一次,先打自己再打子树
print_treeB
for b in branches(tree):
    print(label(tree))
    print_treeB(b)
3 3 1 2 1 3 1 打印次数 = 分支数。叶子一次都不打印(没有分支就不进循环),有 3 个分支的根被打印 3 次
print_treeC
for b in branches(tree):
    print_treeC(b)
print(label(tree))
-1 1 2 3 1 -1 1 3 后序:先把所有子树打完,最后才打自己,所以根的 3 出现在最后一行

这张表值得盯久一点,因为本讲后半段所有生成器递归都建立在同一个直觉上:「在递归调用之前做事」和「在递归调用之后做事」,产生的顺序完全不同。等下把 print 换成 yield,这个直觉一字不改地继续成立——只不过结果不再是打到屏幕上,而是被一个个「让」给调用方。

2. iterable 与 iterator:书和书签

两个定义,先记住字面意思,再看为什么要分开。

  • 可迭代对象(iterable):任何能被顺序处理的对象。换句话说,能放进 for 循环的就是 iterable。例:列表、元组、字典、字符串、range。
  • 迭代器(iterator):提供对某个 iterable 的元素的按序访问。
    • 对 iterable 调用 iter 就能造出一个 iterator。
    • 对 iterator 调用 next 就能取到下一个元素。

课堂上用的比喻很到位:iterable 是书,iterator 是书签。书本身没有「读到哪儿」这个概念,书里的字是死的、可以被很多人同时读;书签才知道位置。同一本书可以夹很多个书签,各自独立地往前走。

>>> lst = [1, 2, 3]        # iterable:一本书
>>> tracker = iter(lst)    # iterator:夹进去的一个书签
>>> next(tracker)
1
>>> next(tracker)
2
>>> next(tracker)
3
>>> next(tracker)
Traceback (most recent call last):
  ...
StopIteration
逐步推演
1 lst = [1, 2, 3]:全局帧里 lst 绑定到一个列表对象。这个对象自己不含任何「位置」信息。
2 iter(lst):Python 新建一个 list_iterator 对象,它内部记着两样东西——指向 lst 那个列表对象的引用、以及一个初始为 0 的下标。名字 tracker 绑定到这个新对象。注意:列表没有被复制。
3 第一次 next(tracker):读出下标 0 处的元素 1,把内部下标改成 1,返回 1。这一步改变了 tracker 的状态——迭代器是可变的。
4 第二、三次 next 同理,分别返回 2、3,内部下标依次变成 2、3。
5 第四次 next:下标 3 已经越过列表末尾,没有「下一个」可返回了。它不返回 None、不返回 False,而是抛出 StopIteration 异常。
6 这个迭代器从此永久报废。再调 next(tracker) 还是 StopIteration,没有任何办法把它倒回去。想重新读,只能 iter(lst) 造一个新的书签。
对象与状态
Global frame
    lst      ──→ [1, 2, 3]        ← 列表对象(无位置)
    tracker  ──→ list_iterator    ← 迭代器对象
                     .source ──→ 同一个 [1, 2, 3](不是副本)
                     .index  ──→ 0 → 1 → 2 → 3   随每次 next 递增

next(tracker) 第 4 次:.index == 3 已越界 → raise StopIteration

这里的 .source / .index 是我为了讲清楚起的名字,CPython 内部字段不叫这个,你在 Python 里访问不到它们(tracker.index 会报 AttributeError)。但「一个引用 + 一个位置」这个心智模型是准确的。

for 循环其实就是 iter + next + 捕获 StopIteration

现在可以回答第 1 节留下的问题了。for x in lst: ... 这行代码,Python 做的是:

逐步推演
1 求值 in 后面的表达式,得到一个对象。
2 对它调用 iter,拿到一个迭代器。这一步就是 for 要求对象必须是 iterable 的原因——不是 iterable 就没法 iter。
3 反复调用 next:每拿到一个值,就把它绑定到循环变量 x 上,然后执行一遍循环体。
4 某次 next 抛出 StopIteration 时,for 捕获这个异常并正常结束循环(不是报错退出)。

所以 for 从来没有假设过被遍历的东西「是列表」「有 len」「能用下标访问」。它只需要对方能交出一个会 next 的书签。字典、range、文件对象长得千差万别,却都能进 for,原因就在这儿。这也解释了为什么第 1 节的字典能直接遍历:字典的 iter 交出的书签,走的是键。

常见误区

对 iterable 直接调 next。 列表是「书」,不是「书签」,它自己不知道读到哪儿了:

>>> next([1, 2, 3])
Traceback (most recent call last):
  ...
TypeError: 'list' object is not an iterator

报错信息里的 is not an iterator 说得很准:列表是 iterable,但不是 iterator。必须先 iter。

反过来,对不是 iterable 的东西调 iter:

>>> iter(5)
Traceback (most recent call last):
  ...
TypeError: 'int' object is not iterable

这两条报错要能一眼分清:not an iterator = 你少调了一次 iter;not iterable = 这东西根本就不是一串值。

常见误区

把迭代器当列表用。 迭代器只支持一个操作:next。别的都不行:

>>> t = iter([1, 2, 3])
>>> len(t)
Traceback (most recent call last):
  ...
TypeError: object of type 'list_iterator' has no len()
>>> t[0]
Traceback (most recent call last):
  ...
TypeError: 'list_iterator' object is not subscriptable

这不是 Python 偷懒。迭代器压根不知道自己还剩几个元素——本讲后面那个「所有偶数」的生成器有无穷多个元素,len 该返回什么?也不知道第 0 个是什么——它可能早就被 next 消费掉了。

3. 书被改了,书签会怎样

第 2 节强调过:iter(lst) 没有复制 lst。迭代器只是记着「那本书在哪儿」。所以如果有人在你读的过程中改了这本书,你的书签会读到改过之后的内容。课堂上的第二个例子专门演示这件事,它比看上去烧脑:

>>> lst = [1, 3, 2, 7]
>>> iter1 = iter(lst)
>>> next(iter1)
1
>>> lst.append(5)
>>> next(iter1)
3
>>> iter2 = iter(lst)
>>> next(iter2)
1
>>> lst[1] = -10
>>> next(iter2)
-10
逐步推演
1 lst = [1, 3, 2, 7],iter1 = iter(lst):iter1 的位置是 0,指着 lst 这个对象本身。
2 next(iter1) → 读下标 0 得 1,iter1 的位置变成 1。
3 lst.append(5):原地修改(mutation),lst 变成 [1, 3, 2, 7, 5]。iter1 指的还是同一个对象,位置仍是 1,没受影响。
4 next(iter1) → 读下标 1 得 3,位置变成 2。追加发生在末尾,暂时看不出影响——但如果继续 next 下去,iter1 会一路读到那个新加的 5。书变厚了,书签自然会多读几页。
5 iter2 = iter(lst):又夹了一个新书签,位置从 0 开始。iter1 和 iter2 是两个独立的对象,位置互不干扰(此刻 iter1 在 2,iter2 在 0)。
6 next(iter2) → 读下标 0 得 1,iter2 位置变成 1。
7 lst[1] = -10:把下标 1 处的元素换掉,lst 变成 [1, -10, 2, 7, 5]。
8 next(iter2) → 读下标 1。此刻下标 1 上装的已经是 -10 了,于是返回 -10,不是 3。迭代器读的是「读的那一刻」列表里的内容,而不是创建书签时的快照。
对象状态随时间变化
时刻          lst 的内容              iter1.index   iter2.index
─────────────────────────────────────────────────────────────
初始          [1, 3, 2, 7]                0           (不存在)
next(iter1)   [1, 3, 2, 7]                1           (不存在)
append(5)     [1, 3, 2, 7, 5]             1           (不存在)
next(iter1)   [1, 3, 2, 7, 5]             2           (不存在)
iter(lst)     [1, 3, 2, 7, 5]             2                0
next(iter2)   [1, 3, 2, 7, 5]             2                1
lst[1] = -10  [1, -10, 2, 7, 5]           2                1
next(iter2)   [1, -10, 2, 7, 5]           2                2   → 返回 -10
直觉

三句话概括这一节:

  • 迭代器不拷贝数据,它只记位置。数据是共享的。
  • 多个迭代器可以指向同一个 iterable,各走各的。iter1 和 iter2 就是这样。
  • 在迭代过程中修改底层容器,是自找麻烦。这段代码能跑通、结果也可以逐步推出来,但真实项目里没人愿意维护这种代码。上一节字典宁可直接抛 RuntimeError,列表这里连报错都没有,只是安静地给你一个你没预料到的值——后者其实更危险。
注意

迭代器不但记着位置,还随身带着这个位置到处走。这带来一个非常容易踩的后果:

>>> c = iter([1, 2, 3, 4])
>>> next(c)
1
>>> list(c)
[2, 3, 4]
>>> list(c)
[]

第一次 list(c) 只拿到 [2, 3, 4]——因为 1 已经被消费掉了;而这次 list 又把剩下的三个全部消费光,所以第二次 list(c) 得到空列表。这不是 bug,是迭代器的定义。「同一个迭代器遍历两遍,第二遍是空的」是本讲最高频的错因,第 5 节还会以另一种面目再遇到它。

4. iterator vs iterable:包含关系与「为什么要多这一层」

初学者最容易糊涂的一句话:所有 iterator 都是 iterable,但不是所有 iterable 都是 iterator。

Iterators vs. Iterables 的包含关系图
iterator 是 iterable 的真子集。外圈(iterable,「书」)的资格是:能进 for 循环、能被 iter 调用;内圈(iterator,「书签」)额外拥有 next。既然 iterator 在外圈里面,它当然也能进 for 循环、也能被 iter 调用——只不过对 iterator 调 iter 返回的是它自己。

为什么迭代器也算 iterable?因为它满足 iterable 的定义:能放进 for 循环。第 3 节最后 list(c) 那个例子已经证明了这一点——list 内部就是在遍历它。

那对一个已经是迭代器的对象调 iter 会怎样?

>>> a = iter([1, 2, 3])
>>> iter(a) is a
True

返回的就是原对象本身,不是新书签。这条规则看起来是个技术细节,其实很关键:正因为 iter(iterator) 返回自己,第 2 节里 for 的第 2 步(无条件先调一次 iter)才能对迭代器同样适用——for x in some_iterator 不会莫名其妙地从头开始,而是从当前位置继续。

操作对 iterable(如 list)对 iterator(如 list_iterator)
for x in obj可以,每次从头可以,从当前位置继续,且用完即废
iter(obj)返回一个新迭代器,位置为 0返回 obj 自己(iter(a) is a 为 True)
next(obj)TypeError: 'list' object is not an iterator返回下一个元素,或抛 StopIteration
len(obj)可以TypeError: object of type 'list_iterator' has no len()
obj[0]可以(序列类型)TypeError: 'list_iterator' object is not subscriptable
能遍历几次无限次一次

为什么要多这一层?

如果只是想遍历列表,for x in lst 就够了,凭什么还要 iter / next?课上给了三条理由,每一条都值得展开。

理由一:迭代器是加在 iterable 之上的一层抽象(abstraction)。 用迭代器写的代码,对「值是从哪儿来的」不做任何假设:可能来自列表,可能来自文件的每一行,可能来自网络,也可能是现算出来的(第 6 节的生成器)。只要对方能 next,你的代码就能用。于是提供数据的一方可以随时改内部实现——把列表换成数据库查询——而下游代码一行都不用改。这正是这门课反复强调的「用接口隔离实现」,只不过换了个场景。

理由二:迭代器把「元素」和「位置」打包进同一个对象。 这意味着位置可以被传递。看这段代码:

>>> def take_two(it):
...     return [next(it), next(it)]
...
>>> t = iter([10, 20, 30, 40, 50])
>>> take_two(t)
[10, 20]
>>> take_two(t)
[30, 40]
>>> next(t)
50

两次调用 take_two,第二次自动从 30 接着来——因为「读到哪儿了」这个信息装在 t 里,跟着参数一起传进了函数。如果传的是列表,take_two 就必须额外接一个「起始下标」参数,还得想办法把新下标传回来。迭代器让「进度」变成一个可以传参的值。

由此还派生出一个很实用的保证:每个元素只会被处理一次。多个函数轮流从同一个迭代器取值,谁也不会重复处理别人已经处理过的元素。

理由三:迭代器阻止你改动原始的 iterable。 拿到一个迭代器,你能做的只有 next——没有 append、没有下标赋值。把迭代器而不是列表交给别人,就等于说「你只能读,不能改」。(当然,第 3 节说过,如果别人手上另有那个列表的引用,他照样能改;迭代器保护的是通过它本身发起的修改。)

直觉

把这三条串起来:迭代器是一个只读的、一次性的、自带进度的取值接口。「只读」让它安全,「自带进度」让它可传递,「不假设来源」让它通用。接下来两节要讲的所有东西——map/filter、生成器——都是在这个接口上做文章。

5. 内置的惰性函数:map、filter、zip、reversed

Python 有一批内置函数,它们的返回值不是列表,而是迭代器,并且惰性(lazily)计算——需要一个值的时候才算一个,而不是一上来把全部结果算完。

函数产生的元素返回的对象类型
map(func, iterable)对 iterable 里每个 x,产出 func(x)map 对象
filter(func, iterable)只产出使 func(x) 为真的那些 xfilter 对象
zip(first_iter, second_iter)产出同下标配对的元组 (x, y, ...),可以 zip 两个以上zip 对象
reversed(sequence)逆序产出序列中的元素逆序迭代器

它们全都是迭代器,所以第 3、4 节讲的规矩一条不落地适用:一次性、不能 len、不能下标。

>>> map(abs, [-1, 2, -3])
<map object at 0x7ec7e9593ee0>
>>> list(map(abs, [-1, 2, -3]))
[1, 2, 3]

直接在交互式解释器里敲 map(...),看到的是 <map object at 0x...> 而不是结果——这几乎是每个人第一次用 map 时的困惑。它不是出错了,是这个对象还没开始算。

zip 的长度规则

>>> list(zip([1, 2], [3, 4]))
[(1, 3), (2, 4)]
>>> list(zip([1, 2], [3, 4, 5], [6, 7]))
[(1, 3, 6), (2, 4, 7)]

第二个例子里第二个列表有 3 个元素,但结果只有 2 个元组:zip 在最短的那个 iterable 用完时就停,多出来的 5 被丢掉。这个「以最短为准」的规则等下做 palindrome 时会用上。

把迭代器变回能看的东西

想看完整内容,就把元素倒进一个容器:

函数作用
list(iterable)造一个含全部元素的列表
tuple(iterable)造一个含全部元素的元组
sorted(iterable)造一个排好序的列表;和 min/max 一样可以传 key 函数

注意这三个都会把迭代器抽干:它们必须看到每一个元素才能完成工作(sorted 甚至必须看完全部才能排序)。所以一旦 list(t) 过一次,t 就空了。

惰性到底惰到什么程度

课上这个例子是本节的核心,把它一步一步看完,你对「惰性」的理解就到位了。

filter 与 map 嵌套的惰性求值演示
double 每被调用一次就打印一行,于是屏幕上的打印记录直接暴露了「什么时候真的算了」。构造 t 那一行一个字都没打印;直到第一次 next(t),才连续算了 3、4、5 三个数——恰好算到第一个满足 f 的值为止,一个多的都没算。
>>> def double(x):
...     print(f'** {x} => {2 * x} **')
...     return 2 * x
...
>>> f = lambda x: x >= 10
>>> t = filter(f, map(double, range(3, 7)))
>>> next(t)
** 3 => 6 **
** 4 => 8 **
** 5 => 10 **
10
逐步推演
1 range(3, 7) 代表 3、4、5、6 这四个数。
2 map(double, range(3, 7)):什么都没发生。它只是造了一个 map 对象,记着「函数是 double、源头是那个 range」。double 一次都没被调用——屏幕上没有任何打印,这是证据。
3 filter(f, ...):同样什么都没发生,只是造了个 filter 对象,记着「判据是 f、源头是那个 map 对象」。整条 t = ... 语句执行完,屏幕上依然一片空白。
4 第一次 next(t)。filter 要交出一个值,于是向它的源头 map 要一个:map 向 range 要到 3,调用 double(3) → 打印 ** 3 => 6 **,交出 6。filter 检查 f(6) 即 6 >= 10 → False,丢弃,继续要。
5 map 再要到 4,double(4) 打印 ** 4 => 8 ** 并交出 8。f(8) 为假,继续。
6 map 要到 5,double(5) 打印 ** 5 => 10 ** 并交出 10。f(10) 即 10 >= 10 → True。filter 立刻返回 10 并停下。
7 注意 6 此刻还没被 double 处理过。第一次 next 只做了「找到第一个合格值」所必需的最少工作。

继续往下要:

>>> next(t)
** 6 => 12 **
12
>>> list(t)
[]
逐步推演
1 第二次 next(t):从上次停下的地方继续。map 从 range 要到 6,double(6) 打印 ** 6 => 12 ** 交出 12;f(12) 为真,返回 12。
2 list(t):filter 继续向 map 要值,map 向 range 要——range 已经给完了 3、4、5、6,抛 StopIteration;map 把它传下去,filter 也抛。list 捕获后结束,得到 []。
3 结果是空列表,不是 [10, 12]。因为 10 和 12 早就被前两次 next 消费掉了。这就是第 3 节那条「同一个迭代器遍历两遍,第二遍是空的」在真实场景里的样子。
核心结论

惰性求值带来两个实际好处:

  • 省计算:只想要第一个结果,就只付第一个结果的代价。上面 double 被调了 3 次而不是 4 次。
  • 省内存:中间结果不必同时存在。filter(f, map(double, range(10 ** 9))) 这行代码瞬间完成,而 [double(x) for x in range(10 ** 9)] 会把机器撑爆。

代价是:结果只能看一遍,而且要用 list 包一层才能打印出来。

常见误区

把同一个 map/filter 对象用两次。 这是本讲最常见的 bug,长这样:

>>> s = [-4, -3, -2, 3, 2, 4]
>>> mp = map(abs, s)
>>> min(mp)
2
>>> list(mp)          # 想再看看都有哪些绝对值
[]

min 为了找最小值必须看完全部元素,mp 就此被抽干。不报错,只是安静地给你一个空列表——比报错难查得多。要用两次,就写成 abs_vals = list(map(abs, s)),先落成列表再说。

注意

reversed 要求参数是序列(有长度、能按下标倒着取),迭代器不行:

>>> reversed(map(abs, [1, -2]))
Traceback (most recent call last):
  ...
TypeError: 'map' object is not reversible
>>> reversed(iter([1, 2, 3]))
Traceback (most recent call last):
  ...
TypeError: 'list_iterator' object is not reversible

想想也合理:要倒着给你,就得先知道最后一个在哪儿;而迭代器连自己还剩几个都不知道。reversed(range(3)) 倒是可以(range 是序列),得到的是一个 range_iterator,list 之后是 [2, 1, 0]。

6. 生成器:用写函数的方式造迭代器

map 和 filter 好用,但它们能表达的模式很窄——「每个都变换一下」「筛掉一些」。如果我想要的序列是「先 x 再 -x」,或者「树里所有和为 total 的路径」,怎么办?

手写一个迭代器类是可以的(那要等学完面向对象),但 Python 给了一条捷径:把 return 换成 yield,函数就变成了迭代器。

  • 生成器函数(generator function):函数体里出现了 yield 的函数。
  • 生成器对象(generator object):调用生成器函数得到的返回值。可以对它调 next,按顺序拿到被 yield 出来的值。
注意

课上专门提醒过:题目里这两个词有时会被混用,虽然它们不是一回事。看不出题面指的是哪个,就在考试时问监考、在小测里报 issue。自己写代码时最好分清楚——「生成器函数」是那段 def,「生成器对象」是调用它得到的东西。

>>> def plus_minus(x):
...     yield x
...     yield -x
...
>>> plus_minus(5)
<generator object plus_minus at 0x78d032746420>
>>> gen = plus_minus(5)
>>> next(gen)
5
>>> next(gen)
-5
>>> next(gen)
Traceback (most recent call last):
  ...
StopIteration

先看第一行输出:plus_minus(5) 的值不是 5,也不是 -5,而是一个生成器对象。这说明调用生成器函数时发生的事,跟调用普通函数完全不同。

调用生成器函数时,函数体一行都不执行

这是本节最反直觉、也最重要的一点。用带副作用的例子来验证:

>>> def noisy():
...     print('start')
...     yield 1
...     print('middle')
...     yield 2
...     print('end')
...
>>> h = noisy()
>>> # 什么都没打印!
>>> next(h)
start
1
>>> next(h)
middle
2
>>> next(h)
end
Traceback (most recent call last):
  ...
StopIteration
逐步推演
1 h = noisy():Python 看到 noisy 的函数体里有 yield,于是不执行函数体,只创建一个帧并把它挂起,连同「下一句该从函数体第一行开始」这个信息一起打包成生成器对象返回。print('start') 没有执行,所以屏幕上什么都没有。
2 第一次 next(h):那个挂起的帧被恢复,从函数体第一行开始执行。打印 start,走到 yield 1。
3 yield 1 做两件事:把 1 交给 next 的调用方作为返回值;把当前帧原地冻结——局部变量全部保留,执行位置记在 yield 1 这一行。函数没有返回、没有结束,只是暂停。
4 第二次 next(h):帧被解冻,从 yield 1 的下一句继续。打印 middle,走到 yield 2,交出 2,再次冻结。
5 第三次 next(h):解冻,打印 end,然后函数体走到了尽头。此时生成器自动抛出 StopIteration——注意 end 是打印出来了的,异常发生在它之后。
6 从此 h 报废,再 next 永远是 StopIteration。
帧的挂起与恢复
h = noisy()
  ┌─ 创建帧 f1: noisy [parent=Global],但立刻挂起 ─┐
  │   局部变量:(无)                              │
  │   下一条待执行:print('start')                  │
  └────────────────────────────────────────────────┘
  Global frame: h ──→ generator object(内部持有上面这一帧)

next(h) 第 1 次:恢复 f1 → print('start') → 到 yield 1
  f1 冻结,下一条待执行 = print('middle');返回值 1 交给调用方

next(h) 第 2 次:恢复 f1 → print('middle') → 到 yield 2
  f1 冻结,下一条待执行 = print('end');返回值 2 交给调用方

next(h) 第 3 次:恢复 f1 → print('end') → 函数体结束
  f1 销毁 → raise StopIteration

对比一下你熟悉的普通函数:普通函数的帧在 return 时就消失了,局部变量随之作废;生成器的帧则在 yield 处存活着被冻起来,等下一次 next 把它唤醒。「函数可以暂停并保留现场」是生成器唯一的新东西,其余全是它的推论。

对比项普通函数生成器函数
调用时立刻执行函数体不执行任何函数体代码,返回生成器对象
返回值return 后面那个值(没写就是 None)一个生成器对象
能返回几个值一个(return 一执行函数就结束)任意多个,一次 next 交出一个
帧的命运return 时销毁yield 时冻结保留,下次 next 恢复
结束方式返回值给调用方函数体跑完 → 抛 StopIteration
能否放进 for不能能——生成器对象是迭代器

既然生成器对象是迭代器,第 4 节那条规则也适用:

>>> g = plus_minus(5)
>>> iter(g) is g
True
>>> list(plus_minus(5))
[5, -5]
>>> for v in plus_minus(5):
...     print(v)
...
5
-5
常见误区

误区一:对生成器函数调 next。

>>> next(plus_minus)
Traceback (most recent call last):
  ...
TypeError: 'function' object is not an iterator

plus_minus 是函数对象,plus_minus(5) 才是生成器对象。少写一对括号就是这条报错。

误区二:以为调用了生成器函数就等于跑了它。 下面这行代码什么都不做:

>>> def f():
...     print('side effect')
...     yield 1
...
>>> f()
<generator object f at 0x...>

side effect 没有被打印。如果你把生成器当成「一个会打印东西的函数」来用,会以为程序坏了。必须有人 next 它、for 它、或者 list 它,函数体才会真的跑。

误区三:在生成器里写 return <值>,以为能拿到这个值。

>>> def g2():
...     yield 1
...     return 99
...
>>> list(g2())
[1]

99 不在结果里。在生成器函数中,return 的作用是提前结束(等价于函数体跑完),返回值被塞进 StopIteration 异常里,普通的 for / list 看不到它。所以生成器里只写光秃秃的 return 用来提前收工,别指望它传值。

误区四:把 yield 写在函数外面。

    yield 5
    ^^^^^^^
SyntaxError: 'yield' outside function

这是 SyntaxError,整个文件一行都跑不了。

7. 无限生成器:列表做不到的事

生成器最能体现价值的地方,是表示没有尽头的序列。因为它每次只算一个值,「总共有多少个」根本不影响它能不能被创建。

>>> def evens():
...     i = 0
...     while True:
...         yield i
...         i += 2
...
>>> gen = evens()
>>> next(gen)
0
>>> next(gen)
2

while True 在普通函数里是死循环,是 bug;在生成器里它是特性。原因就是第 6 节那条:yield 会把帧冻住,控制权交回调用方。循环是被「暂停」的,不是在空转。

逐步推演
1 gen = evens():建帧、挂起。i 此刻还不存在(i = 0 都没执行)。
2 next(gen):恢复,执行 i = 0,进入 while True,条件为真,执行 yield i → 交出 0,冻结。此刻帧里 i 是 0,下一条待执行是 i += 2。
3 next(gen):恢复,执行 i += 2 → i 变成 2;回到 while 条件,仍为真;执行 yield i → 交出 2,冻结。
4 之后每次 next 都重复第 3 步,i 一直是上次冻结时的值加 2。局部变量 i 跨越多次 next 一直活着,这正是帧被保留的直接后果。
帧在多次 next 之间保持存活
f1: evens [parent=Global]   (被 gen 这个生成器对象持有)

  next 第 1 次结束时:  i ──→ 0    暂停在 yield i
  next 第 2 次结束时:  i ──→ 2    暂停在 yield i
  next 第 3 次结束时:  i ──→ 4    暂停在 yield i
  ...

对比普通函数:每次调用都是一个全新的帧,i 每次都从头开始。
常见误区

对无限生成器调 list。

>>> gen = evens()
>>> list(gen)
(无限循环,程序卡死;Ctrl+C 打断会看到 KeyboardInterrupt,
 放任不管则最终 MemoryError)

list 的工作是「一直 next 到 StopIteration 为止」,而 evens() 永远不会 StopIteration。sorted、min、max、sum、tuple 同理,全都会卡死。

正确的取法是只取有限个:

>>> gen = evens()
>>> [next(gen) for _ in range(5)]
[0, 2, 4, 6, 8]

或者在 for 里配合 break。这也是判断题里的常客:看到 list(某无限生成器),答案就是「卡死」,不是某个列表。

直觉

无限生成器把「序列」和「存储」彻底解耦了。evens() 这个对象在内存里只占几十字节——它存的不是无穷多个偶数,而是一条产生偶数的规则加上当前进度。这跟数学里写 $a_n = 2n$ 而不是把所有 $a_n$ 列出来是同一种想法。

8. yield from:把一整串值转手让出去

经常会遇到这种需求:我这个生成器要产出的值,一部分来自另一个 iterable。比如「先把 a 里的元素全给出去,再把 b 里的全给出去」:

def a_then_b(a, b):
    for x in a:
        yield x
    for x in b:
        yield x

这段代码没错,但 for x in a: yield x 这个模式太常见了,Python 给了专门的语法:

def a_then_b(a, b):
    yield from a
    yield from b
>>> list(a_then_b([1, 2], [3, 4]))
[1, 2, 3, 4]

yield from <iterable> 的意思是:把这个 iterable 里的元素一个一个地 yield 出去,直到它耗尽。它不是「yield 一个 iterable」——这个区别下面就要吃它的亏。

核心结论

yield x:让出一个值 x。
yield from xs:让出 xs 里的每一个值,相当于 for v in xs: yield v。
所以 yield from 后面必须跟一个 iterable(列表、字符串、另一个生成器都行),跟一个整数会报 TypeError: 'int' object is not iterable。

递归生成器

课上说得明白:用 yield from 的一个常见理由,就是递归地调用生成器函数。 看这个倒计时:

def countdown(k):
    if k > 0:
        yield k
        yield from countdown(k - 1)
    else:
        yield 'Blastoff!'
>>> list(countdown(3))
[3, 2, 1, 'Blastoff!']

这个结构和上一讲所有树递归、和第 1 节的 print_treeA 是同一个骨架:先处理自己(yield k),再把「更小的同类问题」的全部结果转手让出去(yield from countdown(k - 1))。把它真的展开一层层看:

逐步推演:list(countdown(3))
1 countdown(3) 创建生成器 G3(不执行函数体)。list 开始反复 next(G3)。
2 第 1 次 next:G3 恢复,3 > 0 为真,执行 yield 3 → 产出 3,G3 冻结在这一行。
3 第 2 次 next:G3 恢复,执行 yield from countdown(2)。这一句先调用 countdown(2) 得到生成器 G2,然后 G3 把控制权转给 G2:next(G2) → G2 里 2 > 0 为真,yield 2 → 产出 2。这个 2 穿过 G3 直达 list。此刻 G3 冻在 yield from 那一行,G2 冻在 yield 2。
4 第 3 次 next:G3 恢复 → 它还在 yield from G2 里,于是 next(G2) → G2 执行 yield from countdown(1),创建 G1,next(G1) → yield 1 → 产出 1。此刻三层帧全部活着:G3 冻在 yield from,G2 冻在 yield from,G1 冻在 yield 1。
5 第 4 次 next:层层往下 → G1 恢复,执行 yield from countdown(0),创建 G0,next(G0) → 0 > 0 为假,走 else,yield 'Blastoff!' → 产出 'Blastoff!'。这就是 base case。
6 第 5 次 next:G0 恢复,函数体到头 → StopIteration。G1 的 yield from 捕获它、认为 G0 耗尽,于是 G1 继续往下——函数体也到头了 → G1 抛 StopIteration。同理逐层上传:G2 结束、G3 结束。list 收到 StopIteration,停止。
7 回代:list 收集到的依次是 3、2、1、'Blastoff!',即 [3, 2, 1, 'Blastoff!']。
四次 next 后的帧栈(全部处于冻结状态)
list  ← 'Blastoff!'
 └─ G3: countdown  k=3   暂停在 yield from countdown(k-1)
     └─ G2: countdown  k=2   暂停在 yield from countdown(k-1)
         └─ G1: countdown  k=1   暂停在 yield from countdown(k-1)
             └─ G0: countdown  k=0   暂停在 yield 'Blastoff!'

每一层的 k 都独立保存在自己的帧里,互不干扰。
一个值要从 G0 传到 list,得穿过三层 yield from。

忘了 from 会怎样

用 yield 代替 yield from 的错误版本
把 yield from countdown(k - 1) 写成 yield countdown(k - 1),结果只有两个元素:3 和一个生成器对象。因为 yield 只让出「一个值」,而这个值恰好是生成器对象本身;递归到此为止,里面的 2、1、'Blastoff!' 一个都没被取出来。
def countdown(k):
    if k > 0:
        yield k
        yield countdown(k - 1)      # 少了一个 from
    else:
        yield 'Blastoff!'
>>> list(countdown(3))
[3, <generator object countdown at 0x7f7b63b46420>]
逐步推演:为什么只有两个元素
1 第 1 次 next:yield 3 → 产出 3。
2 第 2 次 next:执行 yield countdown(2)。先求值算子数 countdown(2)——按第 6 节的规则,这只是创建一个生成器对象,函数体一行不跑。然后 yield 把这个对象本身当成一个普通的值让出去。
3 第 3 次 next:函数体已经到头,抛 StopIteration。
4 于是 list 收到的两个元素是 3 和那个从未被 next 过的生成器对象。它里面的 2、1、'Blastoff!' 永远不会被算出来——没人向它要过值。
常见误区

看到输出里冒出 <generator object ...>,99% 是漏了 from。 这个报错很友好,因为它根本不报错——程序照常运行,只是结果里混进了一个你不认识的对象。诊断口诀:

  • 结果里有 <generator object> → 某处该写 yield from 却写了 yield。
  • 反过来,yield from 后面跟了个非 iterable(比如 yield from label(t),而 label 是个整数)→ TypeError: 'int' object is not iterable。

还有一个更隐蔽的变体:想让出一个列表本身作为单个元素时,必须用 yield。第 11 节的 yield_paths 产出的每个元素都是一条路径(一个列表),那里就必须写 yield [label(t)] + path,写成 yield from 会把路径拆成一个个数字。

9. 三个同心圆:iterable、iterator、generator

三个名词到这儿全都出场了,把关系一次理清。

Iterables vs. iterators vs. generators 的同心圆关系
三层同心:generator ⊂ iterator ⊂ iterable。生成器对象是一种迭代器(能 next),迭代器是一种可迭代对象(能进 for)。图上方两条规则说明 iter 的行为:喂给它 iterable 得到新的 iterator,喂给它 iterator 得到的还是它自己。
iterableiteratorgenerator object
典型例子list、tuple、str、dict、rangeiter(lst)、map/filter/zip/reversed 的返回值调用生成器函数的返回值
能进 for✓✓✓
能 next✗✓✓
iter(x) 返回一个新迭代器x 自己x 自己
能重复遍历✓(每次 for 拿新书签)✗✗
元素存在内存里吗通常是(range 除外)不一定不存,用时才算
怎么造出来字面量、构造函数iter(...) 或内置惰性函数写一个含 yield 的函数并调用它

用一句话串起来:iterable 是「有一串值」这个资格;iterator 在此之上加了「记着读到哪儿」;generator 则是「用一段暂停/恢复的函数体来现算这一串值」的迭代器。 每往里一层,能力更强,但也更「一次性」。

注意

「生成器函数」不在这三个圈里的任何一个。它是一个普通的函数对象,只不过调用它会生产出圈内的东西。plus_minus 不是 iterable(for x in plus_minus 会报 TypeError: 'function' object is not iterable),plus_minus(5) 才是。

10. 案例研究:从 count_partitions 到 yield_partitions

这一节把整讲的东西合起来用一次。题目是 Lecture 06 见过的整数分拆(partition):把正整数 n 写成若干个不超过 m 的正整数之和,各部分按递增顺序排列,问有多少种写法。n = 6, m = 4 时有 9 种:

6 = 2 + 4        6 = 1 + 1 + 4      6 = 3 + 3
6 = 1 + 2 + 3    6 = 1 + 1 + 1 + 3  6 = 2 + 2 + 2
6 = 1 + 1 + 2 + 2                   6 = 1 + 1 + 1 + 1 + 2
6 = 1 + 1 + 1 + 1 + 1 + 1

第一步:原始的计数版本

def count_partitions_original(n, m):
    if n == 0:
        return 1
    elif n < 0:
        return 0
    elif m == 0:
        return 0
    else:
        with_m = count_partitions_original(n - m, m)
        without_m = count_partitions_original(n, m - 1)
        return with_m + without_m

核心的拆分思路是对「最大的那部分用不用 m」做二分:要么这个分法里至少用一次 m(剩下要凑 n - m,还能继续用 m),要么完全不用 m(凑 n,但只能用到 m - 1)。两类不重不漏,加起来就是全部。

第二步:合并 base case,为改写做准备

随堂代码给了一个等价的简化版,它把「n < 0」和「m == 0」并成一个失败出口,并把「n == 0」这个成功出口挪进 else:

def count_partitions_simplified(n, m):
    if n < 0 or m == 0:
        return 0
    else:
        exact_match = 0
        if n == m:
            exact_match = 1
        with_m = count_partitions_simplified(n - m, m)
        without_m = count_partitions_simplified(n, m - 1)
        return exact_match + with_m + without_m

为什么这么改?因为 n == 0 那个 base case 的含义是「刚好凑满了,算一种分法」,但它不知道自己是怎么凑满的——返回 1 就完了。等下我们要返回的不是数量而是分法本身,就必须在还知道最后一块是什么的时候记下来。n == m 正是那个时刻:还剩 n 要凑,而允许的最大部分恰好是 m == n,那就直接用一块 m 填满,这是一种分法。

第三步:返回列表而不是数量

把每个 return <数> 换成 return <列表>,加法换成列表拼接:

def list_partitions(n, m):
    if n < 0 or m == 0:
        return []
    else:
        exact_match = []
        if n == m:
            exact_match = [[m]]
        with_m = [p + [m] for p in list_partitions(n - m, m)]
        without_m = list_partitions(n, m - 1)
        return exact_match + with_m + without_m
计数版列表版为什么这样对应
return 0return []「0 种分法」= 「空的分法列表」
exact_match = 1exact_match = [[m]]那一种分法就是 [m];外面再套一层,因为返回的是分法的列表
count(n - m, m)[p + [m] for p in list_partitions(n - m, m)]子问题的每个分法都要补上刚用掉的那个 m;加在末尾是为了保持递增顺序
+(数相加)+(列表拼接)把三部分结果并起来

随堂代码里还有一个 list_partitions_str,除了把 p + [m] 换成 f"{p} + {m}"、把 [[m]] 换成 [str(m)],其余一模一样,输出形如 '1 + 2 + 3'。

第四步:改成生成器

列表版有个硬伤:它必须把全部分法算完才返回。list_partitions(60, 50) 要生成 966370 个分法,在我的机器上要跑约 3 秒,而且这 96 万个结果同时占着内存。生成器版可以边算边给:

def yield_partitions(n, m):
    if n > 0 and m > 0:
        if n == m:
            yield str(m)
        for p in yield_partitions(n - m, m):
            yield f"{p} + {m}"
        yield from yield_partitions(n, m - 1)

逐行对照列表版:

行做什么为什么是这个写法
if n > 0 and m > 0:失败出口取反生成器不需要 return []——什么都不 yield,自然就是空序列。所以直接把「不失败」的条件包在外面,else 分支可以省掉
if n == m: yield str(m)产出「一块 m 填满」这个分法对应 exact_match = [[m]]。这里用 yield 而非 yield from:str(m) 是一个结果。(顺带一提,yield from str(m) 会把 '12' 拆成 '1'、'2',因为字符串也是 iterable)
for p in yield_partitions(n - m, m): yield f"..."递归拿到子问题的每个分法,补上 + m 再产出这里不能用 yield from——每个 p 都要加工之后才能交出去。这正是 yield from 的适用边界:原样转发用 yield from,要加工就得 for + yield
yield from yield_partitions(n, m - 1)不用 m 的那一支,结果原样转发不需要任何加工,所以 yield from 正合适
逐步推演:yield_partitions(6, 4) 的第一个值是怎么冒出来的
1 gen = yield_partitions(6, 4):建帧、挂起,什么都没算。
2 next(gen):恢复。6 > 0 and 4 > 0 为真。n == m?6 != 4,跳过。
3 进入 for p in yield_partitions(2, 4),向子生成器要第一个 p。
4 子生成器 (2, 4):2 != 4,跳过 exact_match;进入 for p in yield_partitions(-2, 4) —— -2 > 0 为假,这个生成器一个值都不产出,for 直接结束(这就是失败出口,无需写 return []);接着 yield from yield_partitions(2, 3)。
5 (2, 3):2 != 3;for p in yield_partitions(-1, 3) 空转;yield from yield_partitions(2, 2)。
6 (2, 2):n == m 命中!yield str(2) → 产出 '2'。这个值原样穿过 (2,3)、(2,4) 两层 yield from,回到第 3 步那个 for,绑定给 p。
7 回到最外层:yield f"{p} + {m}" = yield '2 + 4' → 第一个值是 '2 + 4',也就是 6 = 2 + 4。此时所有帧冻结,剩下 8 个分法一个都没算。
>>> gen = yield_partitions(6, 4)
>>> next(gen)
'2 + 4'
>>> next(gen)
'1 + 1 + 4'
>>> list(yield_partitions(6, 4))
['2 + 4', '1 + 1 + 4', '3 + 3', '1 + 2 + 3', '1 + 1 + 1 + 3', '2 + 2 + 2', '1 + 1 + 2 + 2', '1 + 1 + 1 + 1 + 2', '1 + 1 + 1 + 1 + 1 + 1']

9 个,和计数版的答案对上了。

惰性到底快多少

>>> s = list(yield_partitions(60, 50))   # 要等几秒
>>> len(s)
966370
>>> gen = yield_partitions(60, 50)       # 瞬间完成
>>> next(gen)
'10 + 50'
>>> next(gen)
'1 + 9 + 50'

两种写法算的是同一件事,差别只在什么时候算。list(...) 把 966370 个分法全算出来(约 3 秒),只为了让你看第一个;next(gen) 只做了「找出第一个分法」所必需的那几十次递归调用,感觉不到耗时。如果你的程序只需要前几个结果、或者会在中途 break,生成器版省下的就是几乎全部的计算。

直觉

这四个版本值得整体回看一遍:递归的骨架从头到尾没变过——永远是「用一次 m」+「不用 m」+「恰好填满」这三支。变的只有「结果怎么攒起来」:数量用 +,列表用拼接,生成器用 yield。这就是这门课反复训练的能力:先把问题的递归结构想清楚,再决定用什么形式承载答案。

11. 随堂练习:palindrome 与 min_abs_indices

起始代码是课程网站 Lecture 10 下的 10.py,用 python3 -m doctest 10.py 跑测试(没有输出就是全过)。

palindrome:判断回文

def palindrome(s) -> bool:
    """
    Return `True` if a sequence `s` is the same forward and backward
    (e.g. if `s` is a palindrome), or `False` otherwise.

    >>> palindrome([3, 1, 4, 1, 5])
    False
    >>> palindrome([3, 1, 4, 1, 3])
    True
    >>> palindrome('seveneves')
    True
    >>> palindrome('seven eves')
    False
    """

题目要什么。 判断序列正着读和倒着读是否一样。注意 doctest 里既有列表又有字符串,所以实现不能只对其中一种奏效。最后一个例子 'seven eves' 说明空格算一个字符——不要自作聪明去掉空格。

怎么想到的。 直觉是「把 s 倒过来,和 s 比一比」。倒过来用 reversed(s)。于是第一版:

return reversed(s) == s        # 错的

拿 'aba' 试一下,得到 False——明明是回文。原因是 reversed(s) 返回的是一个迭代器对象,拿一个迭代器去和字符串比较,比的是「这两个是不是同一个对象」,永远是 False。第 5 节强调过:惰性函数的返回值不是你想要的那串值,必须先落进容器。改成:

return list(reversed(s)) == list(s)

两边都套 list,把类型统一成列表再比。右边的 list(s) 不能省:s 是字符串时,list(reversed('aba')) 是 ['a', 'b', 'a'],而 'aba' 是字符串,两者永不相等。

第二种写法(用 zip)。 挑战要求用 zip 和 reversed。思路换成「逐位对比」:把正序和倒序配对,如果每一对的两个元素都相等,就是回文。

return all([a == b for a, b in zip(s, reversed(s))])
逐步推演:palindrome('aba')
1 reversed('aba') 产出 'a'、'b'、'a'(倒序)。
2 zip('aba', reversed('aba')) 把两者同下标配对:('a','a')、('b','b')、('a','a')。注意 zip 也是惰性的,此刻还没真的配。
3 列表推导式遍历这些元组,a, b = ... 是解包,得到 [True, True, True]。遍历这一步把 zip 抽干了,但没关系——结果已经落成列表。
4 all([True, True, True]) → True。(all 的含义是「全部为真」,空列表也返回 True,正好符合「空序列是回文」。)
5 换成 'ab':配对是 ('a','b')、('b','a'),推导式得到 [False, False],all 返回 False。

这个写法其实做了双倍的比较——第 i 位和倒数第 i 位比了两次。但它对,而且很直观。

常见误区
  • reversed(s) == s:比的是迭代器对象和序列,恒为 False。
  • list(reversed(s)) == s:字符串输入时恒为 False(列表 ≠ 字符串)。
  • 用 s[::-1] == s(切片反转)对列表和字符串都能跑,但这题的挑战明确要求用 reversed——而且 s[::-1] 对一般的 iterable 不成立。
  • 把 reversed(s) 存进变量后用两次:r = reversed(s); list(r) == list(s) and all(...r...),第二次 r 已经空了。

min_abs_indices:找出绝对值最小的所有下标

def min_abs_indices(s) -> list[int]:
    """
    Returns a list of all indices of elements in s
    whose absolute value is equal to the minimum absolute value of s.

    >>> min_abs_indices([-4, -3, -2, 3, 2, 4])
    [2, 4]
    >>> min_abs_indices([1, 2, 3, 4, 5])
    [0]
    """

题目要什么。 先算出 s 中所有元素绝对值的最小值,再返回所有绝对值等于它的元素的下标。两个容易读漏的点:返回的是下标不是元素值;可能有多个(第一个 doctest 里 -2 和 2 绝对值都是 2,下标 2 和 4 都要)。

怎么想到的。 拆成两步就清楚了:

1 求最小绝对值。「对每个元素取绝对值」正是 map(abs, s),再喂给 min:min_abs = min(map(abs, s))。
2 挑出下标。既然要的是下标,就遍历下标而不是遍历元素——for i in range(len(s)),保留满足 abs(s[i]) == min_abs 的那些 i。

第 2 步如果直接 for x in s,你会拿到元素却不知道它在第几位,只能再想办法找回下标——那是弯路。要什么就遍历什么。

def min_abs_indices(s):
    min_abs = min(map(abs, s))
    return [i for i in range(len(s)) if abs(s[i]) == min_abs]

第二种写法(不用列表推导式)。 挑战要求只用迭代内置函数。「从一堆候选里挑出满足条件的」就是 filter;候选是所有下标,即 range(len(s)):

def min_abs_indices(s):
    min_abs = min(map(abs, s))
    f = lambda i: abs(s[i]) == min_abs
    return list(filter(f, range(len(s))))

两点值得注意:f 是一个闭包,它用到了外层的 s 和 min_abs——这是第 3 讲高阶函数的老本行,只是这次被 filter 调用;末尾的 list(...) 不能省,因为 filter 返回的是迭代器,而 doctest 期望的输出是 [2, 4],不是 <filter object at 0x...>。

逐步推演:min_abs_indices([-4, -3, -2, 3, 2, 4])
1 map(abs, s) 惰性产出 4, 3, 2, 3, 2, 4;min 把它抽干,得到 min_abs = 2。
2 range(len(s)) 即 range(6),候选下标 0..5。
3 逐个检查:abs(s[0])=4 ✗;abs(s[1])=3 ✗;abs(s[2])=abs(-2)=2 ✓;abs(s[3])=3 ✗;abs(s[4])=abs(2)=2 ✓;abs(s[5])=4 ✗。
4 结果 [2, 4],与 doctest 一致。第二个 doctest [1,2,3,4,5]:最小绝对值 1,只有下标 0 满足,返回 [0]。
常见误区

把 map 对象存起来用两次——第 5 节那个坑在这题里最容易复现:

abs_vals = map(abs, s)
min_abs = min(abs_vals)
return [i for i in range(len(s)) if abs_vals[i] == min_abs]

两处都错:abs_vals 已被 min 抽干;而且 map 对象不支持下标,会报 TypeError: 'map' object is not subscriptable。要么像标准解法那样直接写 abs(s[i]),要么一开始就 abs_vals = list(map(abs, s))。

另一个常见错误是返回元素而不是下标:[x for x in s if abs(x) == min_abs] 会得到 [-2, 2]。读 doctest 的时候就该发现输出是 [2, 4]——两个正数,而原列表里根本没有 4 以外的正 4,说明它们是位置不是值。

12. 随堂练习:yield_paths(递归生成器)

def yield_paths(t, total):
    """
    Yield each path from the root node to any other node in the tree
    that adds up to the target total.
    >>> t = tree(3, [tree(-1), tree(1, [tree(2, [tree(1)]), tree(3)]), tree(1, [tree(-1)])])
    >>> list(yield_paths(t, 3))  # path does not have to go to a leaf
    [[3], [3, 1, -1]]
    >>> list(yield_paths(t, 4))
    [[3, 1], [3, 1]]
    >>> list(yield_paths(t, 5))
    []
    >>> list(yield_paths(t, 6))
    [[3, 1, 2]]
    >>> list(yield_paths(t, 7))
    [[3, 1, 2, 1], [3, 1, 3]]
    """
Practice: Yield Paths 与 doctest 中那棵树的图示
右边就是 doctest 里那棵 t:根 3 有三个孩子 -1、1、1;中间那个 1 有孩子 2 和 3,而 2 又有孩子 1;右边那个 1 有孩子 -1。从根出发沿箭头走到任意节点(不必走到叶子),把沿途标签加起来,等于 total 的路径就要被 yield 出来。

题目要什么。 三个关键点全藏在 doctest 里:

  • list(yield_paths(t, 3)) 的结果里有 [3]——只含根节点的路径也算,而且那句注释明说了 path does not have to go to a leaf,路径可以在任何节点停下。
  • list(yield_paths(t, 4)) 得到 [[3, 1], [3, 1]]——两条一模一样的路径。因为根有两个标签为 1 的孩子,它们是不同的节点,各算一条。所以不能去重。
  • list(yield_paths(t, 5)) 是 []——一条都没有时,什么都不 yield,list 得到空列表。生成器天然支持这个,不需要特判。

怎么想到的。 提示说「参考上一讲 count_paths」。那个函数是数有多少条和为 total 的路径:

def count_paths(t, total):
    if label(t) == total:
        found = 1
    else:
        found = 0
    return found + sum([count_paths(b, total - label(t)) for b in branches(t)])

它的递归结构是:当前节点自己算不算一条(label(t) == total),加上每个分支里、目标减去当前标签之后还有多少条。第 10 节刚练过同一套改造手法——把「数量」换成「东西本身」:

count_pathsyield_paths为什么
found = 1yield [label(t)]「找到一条」= 「产出一条路径」。这条路径只有当前节点自己,所以是 [label(t)]
found = 0(else 分支)什么都不写「0 条」= 「不 yield」。生成器不需要 else
count_paths(b, total - label(t))for path in yield_paths(b, total - label(t))子树里的每条路径,都要拿出来加工
sum([...])yield [label(t)] + path子树给的 path 是从 b 出发的,要变成「从 t 出发」就得在前面接上当前标签

剩下要想明白的只有一件事:为什么目标是 total - label(t)? 因为路径必须从根出发,所以当前节点的标签是一定会被算进去的。既然它已经贡献了 label(t),那么子树里的那一段就只需要凑出 total - label(t)。

def yield_paths(t, total):
    if label(t) == total:
        yield [label(t)]

    for b in branches(t):
        for path in yield_paths(b, total - label(t)):
            yield [label(t)] + path

代码逐行讲

1 if label(t) == total: yield [label(t)] —— 「路径就到我为止」这一种情况。注意这里没有 return:yield 完之后代码继续往下走,因为即使当前节点自己就凑够了,它的子树里可能还有更长的路径也凑够(total=3 时 [3] 和 [3, 1, -1] 就同时存在,后者靠 1 + (-1) = 0 凑回来)。写了 return 就会漏掉 [3, 1, -1]。
2 for b in branches(t) —— 每个分支都要试。叶子节点没有分支,这个循环不执行一次,递归自然终止:这就是 base case,不需要写 if is_leaf(t)。
3 for path in yield_paths(b, total - label(t)) —— 递归调用返回的是一个生成器对象;用 for 遍历它,每次拿到一条「从 b 出发、和为 total - label(t)」的路径。
4 yield [label(t)] + path —— 加工后再让出。这里必须是 yield 而不能是 yield from:我们要交出的是「一条路径」,它本身是一个列表;写成 yield from [label(t)] + path 会把这个列表拆成一个个数字吐出来,结果会变成 [3, 1, -1, ...] 这样的扁平数字流。
5 顺序:[label(t)] + path 是新建一个列表([label(t)] 在前),不是 path.append。用 + 而不是原地修改,避免了多条路径共享同一个列表对象的别名(aliasing)问题。

验证:完整展开 yield_paths(t, 3)

逐步推演
1 顶层 yield_paths(t, 3),label(t) = 3。3 == 3 ✓ → 产出 [3]。这是结果里的第一个。
2 进入 for b in branches(t),新目标一律是 3 - 3 = 0。
3 第一个分支 tree(-1),目标 0。 -1 == 0?否。它是叶子,for b in branches 不执行。这个生成器一个值都不产出,外层的 for path in ... 直接结束。
4 第二个分支 tree(1, [tree(2, [tree(1)]), tree(3)]),目标 0。 1 == 0?否。往下:
· 孙子 tree(2, [tree(1)]),目标 0 - 1 = -1:2 == -1?否 → 曾孙 tree(1),目标 -1 - 2 = -3:1 == -3?否,叶子,无产出。
· 孙子 tree(3),目标 -1:3 == -1?否,叶子,无产出。
整棵子树一无所获。
5 第三个分支 tree(1, [tree(-1)]),目标 0。 1 == 0?否。进入它的分支:孙子 tree(-1),目标 0 - 1 = -1。-1 == -1 ✓ → 孙子产出 [-1]。
6 逐层回代(第一层):这个 [-1] 被第三个分支的 for path 接住,它 yield [1] + [-1] = [1, -1]。
7 逐层回代(第二层):[1, -1] 被顶层的 for path 接住,顶层 yield [3] + [1, -1] = [3, 1, -1]。
8 分支遍历完毕,顶层函数体结束 → StopIteration。list 收集到的是 [[3], [3, 1, -1]],与 doctest 一致 ✓

再看 total = 4 为什么会出现两条相同的 [3, 1]:顶层 3 == 4?否,不产出。三个分支的新目标都是 4 - 3 = 1。第一个分支标签 -1 != 1,无;第二个分支标签 1 == 1 ✓ 产出 [1],回代成 [3, 1];第三个分支标签也是 1 == 1 ✓ 产出 [1],回代成另一个 [3, 1]。两条路径走的是不同的边,只是标签恰好相同——所以结果里有两个 [3, 1],这不是 bug。

常见误区
  • 在第一个 if 里写 return(或 elif/else 把后面的 for 挡住):list(yield_paths(t, 3)) 会变成 [[3]],漏掉 [3, 1, -1]。找到一条路径不代表更长的路径不存在。
  • 递归时忘了减 label(t),写成 yield_paths(b, total):目标一直不变,结果全乱。可以拿 total=6 检验,正确答案是 [[3, 1, 2]](3+1+2=6)。
  • 用 yield from 让出加工后的路径:yield from [label(t)] + path 会把路径拆成散装数字。判断标准很简单——你想让出的是「一个东西」还是「一串东西」。
  • 用 path.insert(0, label(t)) 代替 [label(t)] + path:这是原地修改子生成器给你的那个列表。这题里恰好也能过(每个 path 只被用一次),但一旦上层把同一个列表对象再传给别处就会出诡异 bug。递归里拼接结果,优先用 + 造新对象。
  • 只考虑叶子路径:写成 if is_leaf(t) and label(t) == total。total=3 时 [3](根本身)就会丢掉,doctest 那句注释就是专门防这个的。
直觉

回头看第 1 节的 print_treeA:print(label(tree)) 在前、递归在后。yield_paths 的骨架完全一样——yield 自己这一条在前,遍历分支在后。把树递归里的 print 换成 yield,你就得到了一个能被 for、被 list、被 next 使用的版本,而且不必先把所有结果攒进一个列表。这是生成器在树/图搜索里最典型的用法。

本讲小结

概念要点典型陷阱
iterable(可迭代对象)能进 for 循环的东西。「书」对它调 next → TypeError: 'list' object is not an iterator
iterator(迭代器)iter(iterable) 得到。「书签」,记着位置,只能 next一次性:遍历过就空了,第二次 list 得到 []
iter对 iterable 造新书签;对 iterator 返回它自己以为 iter(it) 能把 it 倒回开头
next取下一个;没有了就抛 StopIteration以为耗尽时返回 None
底层可变性iter 不复制数据,改了原容器会读到新值字典迭代中增删键 → RuntimeError: dictionary changed size during iteration
map/filter/zip/reversed返回惰性迭代器;不 next 就一次都不算直接打印看到 <map object at 0x...>;同一个对象用两次第二次为空
zip同下标配对,以最短的为准以为多出来的元素会配 None
list/tuple/sorted把迭代器抽干,落成容器对无限生成器用 → 死循环 / MemoryError
生成器函数函数体里有 yield;调用它不执行任何函数体代码以为调用了就会打印/计算;对函数本身 next → 'function' object is not an iterator
生成器对象调用生成器函数的返回值,是迭代器;yield 处冻结帧,下次 next 从那里继续以为 yield 像 return 一样结束函数
生成器里的 return提前结束(等价于函数体跑完)以为 return 99 会让 99 出现在结果里
yield from xs把 xs 里的每个元素逐个让出,等价于 for v in xs: yield v递归时漏写 from → 结果里出现 <generator object ...>
什么时候用 yield要让出一个值(哪怕它本身是列表)yield from [label(t)] + path 会把路径拆成散装数字
递归生成器骨架和普通递归一样:自己这份用 yield,子问题原样转发用 yield from,要加工就 for + yield在产出后写 return,挡住了更长的解
失败出口生成器不需要 return []——什么都不 yield 就是空序列习惯性写 else: return [],在生成器里会提前终止函数

三个圈速记

iterable  ⊇  iterator  ⊇  generator object

能进 for  ─────────────────────────────→  三者都能
能 next   ────────────  iterator 起才能  ─→
iter(x)   ──  新书签  ──  返回 x 自己  ──→  返回 x 自己
遍历几次  ──  无限次  ──     一次      ──→      一次

一段代码里最该问自己的三个问题

1 这个东西是「书」还是「书签」? 决定了能不能 next、能不能遍历第二遍。
2 它被消费过了吗? 前面有没有 list/min/sum/for 抽干过它。
3 这一步真的算了吗? 惰性对象在被 next 之前,里面的函数一次都没被调用。

动手练习

练习 1:这段代码打印什么

>>> s = [1, 2, 3, 4]
>>> t = iter(s)
>>> next(t)
>>> u = iter(s)
>>> next(u)
>>> s[1] = 20
>>> next(t)
>>> list(u)
>>> list(t)
看答案

依次是 1、1、20、[20, 3, 4]、[3, 4]。

语句结果t 的位置u 的位置s 的内容
t = iter(s)—0—[1,2,3,4]
next(t)11—[1,2,3,4]
u = iter(s)—10[1,2,3,4]
next(u)111[1,2,3,4]
s[1] = 20—11[1,20,3,4]
next(t)2021[1,20,3,4]
list(u)[20,3,4]24(耗尽)[1,20,3,4]
list(t)[3,4]4(耗尽)4[1,20,3,4]

两个要点:t 和 u 是两个独立的书签,位置互不影响;但它们指着同一个列表,所以 s[1] = 20 之后谁读到下标 1 都得到 20。

练习 2:惰性的代价

下面这段代码运行后,屏幕上打印了几行 calling?result 是什么?

def f(x):
    print('calling')
    return x * x

nums = map(f, [1, 2, 3])
result = [n for n in nums if n > 100]
total = sum(nums)
看答案

打印 3 行 calling,result 是 [],total 是 0。

  • nums = map(f, [1, 2, 3]):0 行。只是造了个 map 对象,f 一次都没被调用。
  • 列表推导式遍历 nums,把它整个抽干:f(1)、f(2)、f(3) 依次被调用,打印 3 行 calling,产出 1、4、9。没有一个大于 100,所以 result 是 []。
  • sum(nums):nums 已经空了,sum 一个元素也拿不到,返回 0。不报错,也不会再打印 calling。

这就是「迭代器只能遍历一次」造成的静默错误:total 本该是 14,实际是 0,程序却一切正常。要修就在第一行写 nums = list(map(f, [1, 2, 3]))。

练习 3:生成器执行顺序

def g():
    print('A')
    yield 1
    print('B')
    yield 2

gen = g()
print('C')
for x in gen:
    print(x)
看答案

输出:C、A、1、B、2。

1 gen = g():建帧后立刻挂起,A 不会被打印。
2 print('C') → 打印 C。所以 C 反而排在 A 前面。
3 for 先 iter(gen)(返回 gen 自己),再 next:帧恢复 → 打印 A → yield 1 交出 1 并冻结 → 循环体 print(x) 打印 1。
4 第二次 next:从 yield 1 之后继续 → 打印 B → 交出 2 → 循环体打印 2。
5 第三次 next:函数体到头 → StopIteration → for 正常结束。

关键是理解控制权在生成器和调用方之间来回切换:函数体不是一口气跑完的,它每 yield 一次就把控制权交回去一次。

练习 4:写一个生成器

实现 evens_up_to(n),按从小到大的顺序 yield 出 0 到 n(含 n)之间的所有偶数。list(evens_up_to(7)) 应得 [0, 2, 4, 6],list(evens_up_to(0)) 应得 [0],list(evens_up_to(-1)) 应得 []。再用它写一个 evens_then_odds(n):先 yield 所有偶数,再 yield 所有奇数。

看答案
def evens_up_to(n):
    i = 0
    while i <= n:
        yield i
        i += 2

def odds_up_to(n):
    i = 1
    while i <= n:
        yield i
        i += 2

def evens_then_odds(n):
    yield from evens_up_to(n)
    yield from odds_up_to(n)

验证:list(evens_up_to(7)):i = 0 → yield 0 → 2 → yield 2 → 4 → yield 4 → 6 → yield 6 → 8,8 <= 7 为假,循环结束,函数体到头 → StopIteration,得 [0, 2, 4, 6] ✓

list(evens_up_to(-1)):一进函数 0 <= -1 就是假,一个 yield 都没执行,直接结束 → [] ✓ 这正是第 10 节说的「生成器的失败出口不需要写任何东西」。

list(evens_then_odds(7)) 得 [0, 2, 4, 6, 1, 3, 5, 7]:两个 yield from 依次把两个子生成器抽干,顺序就是书写顺序。

如果把 evens_then_odds 写成 yield evens_up_to(n),结果会是 [<generator object>, <generator object>]——漏 from 的经典症状。

练习 5:改写成生成器

下面这个函数返回一棵树里所有叶子的标签组成的列表。把它改写成生成器 yield_leaves(t),使 list(yield_leaves(t)) 得到同样的结果。

def leaf_labels(t):
    if is_leaf(t):
        return [label(t)]
    result = []
    for b in branches(t):
        result = result + leaf_labels(b)
    return result
看答案
def yield_leaves(t):
    if is_leaf(t):
        yield label(t)
    for b in branches(t):
        yield from yield_leaves(b)

三处改动,每处都对应本讲的一条规则:

  • return [label(t)] → yield label(t)。返回「装着一个标签的列表」变成「让出一个标签」。注意不再需要那对方括号,因为 yield 一次就代表一个元素。
  • 不需要 result = [] 和累加。生成器的结果是「让出过的东西的序列」,不需要一个容器来攒。
  • result + leaf_labels(b) → yield from yield_leaves(b)。子树给回来的是一串标签,原样转发,所以用 yield from。

为什么原来的 if is_leaf(t): return ... 现在可以不写 else?因为叶子的 branches(t) 是空的,下面那个 for 一次都不执行,效果和 return 一样。这跟 yield_paths 里的处理是同一个道理。

用 t = tree(3, [tree(-1), tree(1, [tree(2, [tree(1)]), tree(3)]), tree(1, [tree(-1)])]) 验证:叶子按先序依次是 -1、1(2 下面那个)、3、-1,所以两个版本都得到 [-1, 1, 3, -1]。

练习 6:判断对错

下面每句话对不对?

  1. iter(iter([1, 2, 3])) 会返回一个全新的迭代器。
  2. 一个生成器对象可以放进 for 循环。
  3. len(map(abs, [-1, 2])) 返回 2。
  4. 生成器函数里必须至少有一个 yield。
  5. list(filter(lambda x: x > 0, iter([1, -2, 3]))) 返回 [1, 3]。
看答案
  1. 错。 对迭代器调 iter 返回它自己,iter(a) is a 为 True。这是「不能把迭代器倒回开头」的技术原因。
  2. 对。 生成器对象是迭代器,迭代器是 iterable,所以能进 for。但只能进一次。
  3. 错。 TypeError: object of type 'map' has no len()。迭代器不知道自己还有几个元素。
  4. 对——这就是生成器函数的定义:函数体里有 yield(或 yield from)的函数。没有 yield 的就是普通函数,调用它会立刻执行函数体。
  5. 对。 filter 的第二个参数可以是任何 iterable,迭代器当然也行;list 把结果抽干得到 [1, 3]。但要注意那个 iter([1, -2, 3]) 已经被消费光了,如果之前存过变量,之后就用不了了。