Python生成器和yield,给你说明白

下文中,yield 就是生成器的关键。

引子:为啥要有生成器?

你碰到过这种情况没?有个巨大的数据集,大到直接把电脑内存给撑爆了。或者,你写了个挺复杂的函数,每次调用都得记住上次的状态,可这函数本身又不大,专门为它写个类吧,又嫌太麻烦。

这种时候,生成器(Generator)和它的好搭档 yield 语句就能帮上大忙。

读完这篇,你就心里有数了:

  • 生成器是啥,咋用?
  • 咋写出生成器函数和生成器表达式?
  • yield 这玩意儿到底是怎么工作的?
  • 一个生成器函数里能写好几个 yield,怎么玩?
  • 生成器的高级玩法:.send().throw().close()
  • 怎么用一堆生成器搭个“数据流水线”
生成器初体验

生成器函数是PEP 255里提出来的,是一种特殊的函数。普通函数 return 完就完了,它不,它返回一个“惰性迭代器”。这东西你能像列表一样循环它。但关键区别是,它不把内容一股脑全放内存里。你知道的,电脑内存是有限的。

例1:盘大文件

处理大文件,比如CSV,是生成器的老本行。CSV文件就是用逗号把数据分成一列一列的。

假设你想数数一个超大CSV文件里有多少行。常规思路(会出问题)的代码大概是这样:

def get_log_lines(filename):
    file = open(filename)
    # 一口气全读完,再按换行符切开
    content = file.read().split("\n")
    return content

log_lines = get_log_lines("huge_server.log")
line_count = 0
for line in log_lines:
    line_count += 1
print(f"总行数: {line_count}")

这段代码会先把整个文件读到内存里。如果文件比内存还大,那结果就是 MemoryError。在你看到这个错误之前,电脑早就卡得像蜗牛爬了。那咋整?看看生成器版本的:

def get_log_lines(filename):
    for line in open(filename, "r"):
        # yield!关键在这里,一行一行往外给
        yield line

log_lines = get_log_lines("huge_server.log")
line_count = 0
for line in log_lines:
    line_count += 1
print(f"总行数: {line_count}")

你看,get_log_lines 变成了一个生成器函数。它打开文件,一行一行读,读一行,yield 一行。这样,内存里永远只装着一行数据,再大的文件也不怕。输出结果一样,但没内存错误。

你甚至还能写得更简洁,用“生成器表达式”(有点像列表推导式):

log_lines = (line for line in open("huge_server.log"))

记住这个核心区别:

  • yield,你得到的是一个生成器对象,可以慢慢取数据。
  • return,你只能拿到文件的第一行。

例2:整一个无限序列

想产生一个无限的数字序列?用 range(5) 只能生成有限个。电脑内存有限,但生成器可以,因为它不是一次性把所有的数都算出来放好,而是你要一个,它给一个。

def generate_numbers():
    n = 1
    while True:
        yield n
        n += 2  # 每次加2,产生奇数序列: 1, 3, 5, ...

# 我们来试试
odd_gen = generate_numbers()
for _ in range(5):
    print(next(odd_gen))

输出:

1
3
5
7
9

这个函数看起来像个死循环,但因为用了 yield,它每次执行到 yield n 就暂停了,把 n 给你。等你下次再问它要(用 next()),它才从上次暂停的地方接着跑,把 n 加2,然后继续循环。用 for 循环去遍历它,它就会一直产生数字,直到你手动停下(比如按 Ctrl+C)。你甚至可以不用 for,直接用 next() 手动控制:

gen = generate_numbers()
print(next(gen))  # 1
print(next(gen))  # 3
print(next(gen))  # 5

例3:当回文探测器

无限序列能玩出不少花样,比如造个回文数检测器。回文数就是正着读反着读都一样的数,像 121。我们先写个函数判断一个数是不是回文:

def is_palindrome_number(num):
    # 一位数不算,比如 1、2、3... 我们跳过
    if num < 10:
        return False
    original = num
    reversed_num = 0
    while num > 0:
        reversed_num = reversed_num * 10 + num % 10
        num = num // 10
    return original == reversed_num

现在,把我们的无限序列生成器 generate_numbers 和这个判断函数结合起来,就能找出所有奇数的回文数:

for num in generate_numbers():
    if is_palindrome_number(num):
        print(num)

输出会是一长串:11, 33, 55, 77, 99, 101, 111, 121, 131… 直到你手动停止。

注:实际干活时,你基本不用自己写无限序列生成器,itertools 模块里的 itertools.count() 又高效又方便。

深入:生成器到底是个啥机制?

前面我们用了两种方法造生成器:生成器函数和生成器表达式。它们看起来像普通函数,但最大的不同就是用了 yield 而不是 return

看回我们那个无限序列的函数:

def generate_numbers():
    n = 1
    while True:
        yield n   # <-- 关键点在这
        n += 2

yield 这地方,就像一个“暂停并返回”的按钮。它把值 n 扔给调用你的地方,但函数本身没退出,而是“冻”在那了。函数里的变量 n 的值、下一步该执行哪行代码(n += 2),所有这些状态都被妥妥地保存着。下次你调用 next(),它就从 yield 后面那行(n += 2)接着运行。

生成器表达式:一行代码搞定

列表推导式是 [n**2 for n in range(5)],生成器表达式就是把方括号换成圆括号:(n**2 for n in range(5))。看个例子:

# 这是一个列表推导式,一口气算完5个数的平方,存成一个列表
squares_list = [x**2 for x in range(5)]
# 这是一个生成器表达式,它只是一个"公式",没真算呢
squares_gen = (x**2 for x in range(5))

print(squares_list)  # 输出: [0, 1, 4, 9, 16]
print(squares_gen)   # 输出: <generator object <genexpr> at 0x...>

性能咋样?内存和速度的权衡

生成器省内存是出了名的。我们测一下大小:

import sys

# 用列表推导式,算1万个数的平方
list_version = [i**2 for i in range(10000)]
print(f"列表大小: {sys.getsizeof(list_version)} 字节")  # 大概 87624 字节

# 用生成器表达式
gen_version = (i**2 for i in range(10000))
print(f"生成器大小: {sys.getsizeof(gen_version)} 字节")   # 才 120 字节左右

列表是生成器的700多倍!但天下没有免费的午餐。生成器慢一点。比如,对1万个数求和:

import cProfile

# 对列表求和 (先算好所有数,再求和)
cProfile.run('sum([i * 2 for i in range(10000)])')
# 结果: 5 function calls in 0.001 seconds

# 对生成器求和 (边算边求和)
cProfile.run('sum((i * 2 for i in range(10000)))')
# 结果: 10005 function calls in 0.003 seconds

生成器慢了大约三倍。所以,如果内存够用,又追求速度,列表推导式更好;如果内存是瓶颈,或者处理的是无穷序列,那必须用生成器。

核心:yield 语句的真面目

yield 的核心就是控制函数流程。调用一个生成器函数,你得到的是一个生成器对象。当你在这个对象上调用 next() 时,函数体内的代码就跑起来,直到撞上 yield

撞上 yield 那一刻,程序就“暂停”了,把 yield 后面的值返回给调用方(这点像 return),但函数的“现场”被完整保留下来了(局部变量、执行位置、异常处理栈等)。等下次再调用 next(),它就从刚才暂停的地方接着跑。

多个 yield 语句

一个生成器函数里可以有多个 yield,它会依次经过它们:

def multi_yield_demo():
    msg = "Hello"
    yield msg
    msg = "World"
    yield msg
    msg = "Goodbye"
    yield msg

demo = multi_yield_demo()
print(next(demo))  # Hello
print(next(demo))  # World
print(next(demo))  # Goodbye
print(next(demo))  # 这里会抛出 StopIteration 异常

生成器,就像所有迭代器一样,会被“耗尽”。你只能从头到尾完整遍历它一次。用 for 循环遍历,循环会自动处理 StopIteration 异常然后退出。如果你手动用 next(),就得自己捕获这个异常。

注:StopIteration 不是错误,它就像迭代器的“哨兵”,告诉你“没货了”。

进阶:生成器的高级玩法

除了 yield,生成器对象还有三个方法:.send().throw().close()。它们让生成器变得更强大。

.send():给生成器喂数据

.send() 可以向生成器内部发送一个值,这个值会成为当前 yield 表达式的结果。这能让生成器变成一个“协程”,既能往外给数据,也能往里收数据。

我们改造一下回文检测的例子。这次,每找到一个回文数,我们就把它的“下一个数量级”的数字喂给生成器,让它从那里继续找。

def palindrome_coroutine():
    num = 0
    while True:
        if is_palindrome_number(num):  # 复用之前的判断函数
            # 这里 yield 出去了一个值,同时也准备接收一个值
            received = (yield num)
            if received is not None:
                num = received
        num += 1

# 使用这个协程
pal_gen = palindrome_coroutine()
# 先启动生成器,让它跑到第一个 yield 处
next(pal_gen)

for i in range(5):
    # 当前回文数
    current_pal = pal_gen.send(None)  # 或者用 next(pal_gen),效果一样
    print(f"找到回文: {current_pal}")
    # 计算下一个数量级,比如当前是 121,下一个是 1000
    next_magnitude = 10 ** len(str(current_pal))
    # 把这个值 send 进去,生成器会从那个数开始继续找
    pal_gen.send(next_magnitude)

这段代码的逻辑是:生成器找到一个回文数后,我们计算出它的下一个数量级(比如从121算出1000),然后通过 .send(next_magnitude) 把它送回生成器。生成器里的 received 变量就拿到了这个1000,然后把它赋值给 num,接着从1000开始继续寻找下一个回文数。

.throw():往生成器里扔异常

你可以在生成器内部抛出一个异常:

def simple_gen():
    try:
        yield 1
        yield 2
    except ValueError:
        print("内部捕获到了ValueError!")
        yield 3
    yield 4

gen = simple_gen()
print(next(gen))  # 1
print(next(gen))  # 2
# 从外部向生成器抛一个异常
print(gen.throw(ValueError))  # 内部会捕获,并输出"内部捕获到了ValueError!",然后 yield 3
print(next(gen))  # 4

这在需要从外部控制生成器内部行为时很有用。

.close():主动停止生成器

这个方法用来提前关闭一个生成器。被关闭的生成器再被迭代时会抛出 StopIteration 异常。

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

counter = infinite_counter()
print(next(counter))  # 1
print(next(counter))  # 2
counter.close()
print(next(counter))  # 这里会抛出 StopIteration 异常
实战:用生成器搭个数据流水线

假设你有一个巨大的CSV文件,记录了公司融资数据(TechCrunch数据集)。你想统计所有A轮融资的总金额。

传统做法是把整个文件读到 pandas DataFrame里再处理。但如果文件巨大,内存可能扛不住。用生成器,我们可以搭一条“数据流水线”,数据像水流一样,从文件源头,经过一道道工序,最后算出结果。每一步都是一个生成器,负责一道工序。

第一步:数据源,一行一行读文件

# 文件路径
data_file = "startup_funding.csv"
# 流水线源头:产生文件的每一行
lines = (line for line in open(data_file))

第二步:数据清洗,把每行按逗号切开,并去掉换行符

# 第二道工序:把一行变成一列值
row_lists = (line.strip().split(',') for line in lines)

第三步:提取列名

CSV的第一行通常是列名,我们把它取出来:

# 从流水线中取出第一行作为列名
column_names = next(row_lists)

第四步:把每行数据和列名组合成字典

# 第三道工序:把列名和每一行的值打包成字典
company_dicts = (dict(zip(column_names, data_row)) for data_row in row_lists)

第五步:过滤数据,只关心A轮融资,并取出金额

# 第四道工序:只保留 round 列是 'a' 的数据,并提取 raisedAmt 列
series_a_funding = (
    int(company['raisedAmt'])
    for company in company_dicts
    if company.get('round') == 'a'
)

第六步:执行流水线,计算结果

到这一步,流水线才真正开始跑起来。之前的步骤都只是定义好了“工序”,还没真正处理数据。我们调用 sum(),数据就会顺着流水线一路流下来,最终算出总和。

total_amount = sum(series_a_funding)
print(f"A轮融资总金额: ${total_amount}")

完整的流水线代码:

data_file = "startup_funding.csv"

lines = (line for line in open(data_file))
row_lists = (line.strip().split(',') for line in lines)
column_names = next(row_lists)
company_dicts = (dict(zip(column_names, data_row)) for data_row in row_lists)
series_a_funding = (
    int(company['raisedAmt'])
    for company in company_dicts
    if company.get('round') == 'a'
)

total = sum(series_a_funding)
print(f"总A轮融资金额: ${total}")

这段代码里,每一步都只是一个生成器表达式。它们共同组成了一个高效的内存友好型数据管道。你可以随意在管道中间插入或修改工序(比如增加数据清洗、格式转换等),非常灵活。

注:生产环境中处理CSV,请使用Python标准库的 csv 模块,它更健壮、功能更全。这里的例子只是为了展示生成器管道的原理。

挑战一下:你能算出A轮融资的平均金额吗?提示:你需要遍历两次,或者自己写一个既算总和又算个数的生成器函数。试试看!

总结

好了,关于生成器,我们就聊到这。你现在应该掌握了:

  • 怎么写和使用生成器函数、生成器表达式。
  • yield 是怎么让函数“暂停”和“恢复”的。
  • 一个函数里可以有好几个 yield
  • 怎么用 .send() 给生成器喂数据,实现双向通信。
  • 怎么用 .throw() 在生成器里触发异常。
  • 怎么用 .close() 关掉一个生成器。
  • 最重要的是,怎么用一堆生成器搭出数据流水线,高效处理大文件。

更多推荐