Python中的生成器(generator)和yield
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()关掉一个生成器。 - 最重要的是,怎么用一堆生成器搭出数据流水线,高效处理大文件。
更多推荐



所有评论(0)