Python 迭代器与生成器深度解析

Python基础 2026-04-20 10
预计阅读时间:31 分钟

Python 迭代器与生成器深度解析

一、迭代的本质:从 for 循环说起

在深入迭代器之前,我们先思考一个基本问题:为什么 for 循环能遍历列表、字符串、字典这些完全不同的对象?

# 这些看似不同的类型,却都能用 for 循环
for item in [1, 2, 3]:           # 列表
    print(item)

for char in "Python":             # 字符串
    print(char)

for key in {'a': 1, 'b': 2}:      # 字典
    print(key)

for line in open('file.txt'):     # 文件对象
    print(line)

答案就是:迭代器协议。任何实现了 __iter__() 和 __next__() 方法的对象,都可以被迭代。

二、迭代器:可迭代对象与迭代器的区别

2.1 核心概念辨析

# 重要区分:可迭代对象 vs 迭代器

# 可迭代对象 (Iterable):实现了 __iter__ 方法,返回一个迭代器
my_list = [1, 2, 3]  # 列表是可迭代对象,但不是迭代器
print(f"列表是可迭代对象吗? {hasattr(my_list, '__iter__')}")  # True
print(f"列表是迭代器吗? {hasattr(my_list, '__next__')}")      # False

# 迭代器 (Iterator):实现了 __iter__ 和 __next__ 方法
iterator = iter(my_list)  # 通过 iter() 获取迭代器
print(f"iterator 是可迭代对象吗? {hasattr(iterator, '__iter__')}")  # True
print(f"iterator 是迭代器吗? {hasattr(iterator, '__next__')}")      # True

# 迭代器的关键特性:一次性消费
print(f"第一次: {next(iterator)}")  # 1
print(f"第二次: {next(iterator)}")  # 2
print(f"第三次: {next(iterator)}")  # 3
# print(next(iterator))  # StopIteration - 迭代结束

# 每次调用 iter() 都会创建新的迭代器
iterator2 = iter(my_list)
print(f"新迭代器从头开始: {next(iterator2)}")  # 1

2.2 深入理解迭代器协议

# 手动实现 for 循环的等价操作
def manual_for_loop(iterable):
    """模拟 for 循环的内部实现"""
    # 1. 获取迭代器
    iterator = iter(iterable)

    # 2. 不断调用 next()
    while True:
        try:
            item = next(iterator)
            print(f"处理元素: {item}")
            # 这里执行 for 循环体
        except StopIteration:
            # 3. 迭代结束
            break

# 测试
manual_for_loop([1, 2, 3])

# 自定义可迭代对象
class CountDown:
    """可迭代对象:倒计时"""

    def __init__(self, start):
        self.start = start

    def __iter__(self):
        """返回一个迭代器"""
        return CountDownIterator(self.start)

class CountDownIterator:
    """对应的迭代器"""

    def __init__(self, start):
        self.current = start

    def __iter__(self):
        """迭代器自身也必须是可迭代的(返回 self)"""
        return self

    def __next__(self):
        """返回下一个值"""
        if self.current < 0:
            raise StopIteration
        value = self.current
        self.current -= 1
        return value

# 使用自定义迭代器
countdown = CountDown(3)
print("第一次迭代:")
for num in countdown:
    print(f"  {num}")

print("第二次迭代(重新开始):")
for num in countdown:  # 每次 for 循环都会调用 __iter__ 获取新的迭代器
    print(f"  {num}")

2.3 将迭代器与可迭代对象合一

class CountDownCombined:
    """既是可迭代对象,又是自己的迭代器"""

    def __init__(self, start):
        self.start = start
        self.reset()  # 初始化迭代状态

    def reset(self):
        """重置迭代状态"""
        self.current = self.start

    def __iter__(self):
        """返回迭代器(就是自身)"""
        self.reset()  # 每次迭代前重置,允许多次迭代
        return self

    def __next__(self):
        if self.current < 0:
            raise StopIteration
        value = self.current
        self.current -= 1
        return value

# 测试
counter = CountDownCombined(3)
print("第一轮迭代:")
for num in counter:
    print(f"  {num}")

print("第二轮迭代:")
for num in counter:
    print(f"  {num}")

# 但要注意:如果忘记在 __iter__ 中重置,第二次迭代会是空的
class BrokenIterator:
    """错误示例:不可重复迭代的迭代器"""
    def __init__(self, data):
        self.data = data
        self.index = 0

    def __iter__(self):
        return self  # 不重置状态

    def __next__(self):
        if self.index >= len(self.data):
            raise StopIteration
        value = self.data[self.index]
        self.index += 1
        return value

broken = BrokenIterator([1, 2, 3])
print("\n第一次迭代:")
for item in broken:
    print(f"  {item}")

print("第二次迭代(空):")
for item in broken:  # index 已经到末尾了
    print(f"  {item}")  # 不会执行

三、生成器:Python 最优雅的特性之一

3.1 生成器函数基础

# 生成器函数:使用 yield 而不是 return
def simple_generator():
    """最简单的生成器函数"""
    print("生成器开始执行")
    yield 1
    print("第一次 yield 之后")
    yield 2
    print("第二次 yield 之后")
    yield 3
    print("生成器结束")

# 调用生成器函数返回生成器对象,不执行函数体
gen = simple_generator()
print(f"生成器对象: {gen}")
print(f"类型: {type(gen)}")

# 每次调用 next() 执行到下一个 yield
print("\n手动迭代:")
print(f"next(gen) -> {next(gen)}")
print(f"next(gen) -> {next(gen)}")
print(f"next(gen) -> {next(gen)}")

try:
    next(gen)  # 超过 yield 数量,抛出 StopIteration
except StopIteration:
    print("生成器已耗尽")

# 生成器可以直接用于 for 循环
print("\n使用 for 循环:")
for value in simple_generator():
    print(f"获得值: {value}")

3.2 生成器的状态保存

def stateful_generator():
    """演示生成器如何保存局部状态"""
    total = 0
    count = 0

    while True:
        # yield 可以接收外部发送的值
        value = yield (total, count)
        if value is None:
            break
        total += value
        count += 1

# 生成器的双向通信
gen = stateful_generator()

# 1. 启动生成器(必须先用 next() 或 send(None) 启动)
print(f"初始状态: {next(gen)}")  # (0, 0)

# 2. 发送数据
print(f"发送 10: {gen.send(10)}")   # (10, 1)
print(f"发送 20: {gen.send(20)}")   # (30, 2)
print(f"发送 30: {gen.send(30)}")   # (60, 3)

# 3. 关闭生成器
try:
    gen.send(None)  # 触发 break
except StopIteration:
    print("生成器正常结束")

# 实战示例:计算移动平均
def moving_average():
    """计算移动平均值的生成器"""
    total = 0
    count = 0
    average = 0

    while True:
        value = yield average
        total += value
        count += 1
        average = total / count

avg_calc = moving_average()
next(avg_calc)  # 启动生成器

# 实时计算平均值
data = [10, 20, 30, 40, 50]
for value in data:
    avg = avg_calc.send(value)
    print(f"加入 {value:2d},当前平均值: {avg:.1f}")

3.3 生成器表达式

# 列表推导式 vs 生成器表达式
# 列表推导式:立即计算,占用内存
list_comp = [x * x for x in range(1000000)]
print(f"列表推导式占用内存: {list_comp.__sizeof__():,} 字节")

# 生成器表达式:延迟计算,内存友好
gen_expr = (x * x for x in range(1000000))
print(f"生成器表达式占用内存: {gen_expr.__sizeof__():,} 字节")

# 生成器表达式可以直接迭代
print("\n生成器表达式示例:")
squares = (x * x for x in range(5))
for square in squares:
    print(f"  {square}")

# 生成器表达式作为函数参数(可以省略外层括号)
total = sum(x * x for x in range(1000000))
print(f"\n平方和: {total}")

# 嵌套生成器表达式
def read_large_file(file_path):
    """逐行处理大文件"""
    with open(file_path, 'r', encoding='utf-8') as f:
        # 链式生成器表达式
        lines = (line.strip() for line in f)
        non_empty_lines = (line for line in lines if line)
        processed_lines = (line.upper() for line in non_empty_lines)

        for line in processed_lines:
            yield line

# 使用示例
# for line in read_large_file('huge_file.txt'):
#     print(line)

四、生成器的高级应用

4.1 yield from:委托给子生成器

# yield from 简化嵌套生成器
def sub_generator(start, end):
    """子生成器"""
    for i in range(start, end):
        yield i

def main_generator_old():
    """不使用 yield from 的老方法"""
    for i in range(3):
        yield i
    for i in sub_generator(3, 6):
        yield i
    for i in range(6, 9):
        yield i

def main_generator_new():
    """使用 yield from 的新方法"""
    yield from range(3)           # 委托给 range 迭代器
    yield from sub_generator(3, 6) # 委托给子生成器
    yield from range(6, 9)

# 验证两者等价
print("老方法:", list(main_generator_old()))
print("新方法:", list(main_generator_new()))

# yield from 的强大之处:双向通信
def writer():
    """接收数据并写入的生成器"""
    result = []
    while True:
        data = yield
        if data is None:
            break
        result.append(data)
    return result  # 返回值会被 yield from 捕获

def reader_wrapper():
    """包装器,演示 yield from 捕获返回值"""
    result = yield from writer()
    print(f"写入器返回的数据: {result}")
    return f"处理了 {len(result)} 条数据"

# 使用
wrapper = reader_wrapper()
next(wrapper)  # 启动

# 发送数据
for i in range(5):
    wrapper.send(f"数据{i}")

# 结束并获取返回值
try:
    wrapper.send(None)
except StopIteration as e:
    print(f"包装器返回值: {e.value}")

# 实战:扁平化嵌套列表
def flatten(nested_list):
    """递归扁平化任意嵌套的列表"""
    for item in nested_list:
        if isinstance(item, (list, tuple)):
            yield from flatten(item)  # 递归委托
        else:
            yield item

nested = [1, [2, [3, 4], 5], 6, [7, 8]]
print(f"\n扁平化结果: {list(flatten(nested))}")

# 对比不使用 yield from 的版本
def flatten_without_yield_from(nested_list):
    """不使用 yield from 的版本"""
    for item in nested_list:
        if isinstance(item, (list, tuple)):
            for sub_item in flatten_without_yield_from(item):
                yield sub_item
        else:
            yield item

4.2 生成器实现协程

import time
from collections import deque

# 使用生成器实现简单的任务调度器
class TaskScheduler:
    """基于生成器的简单任务调度器"""

    def __init__(self):
        self.tasks = deque()

    def add_task(self, task):
        """添加任务"""
        self.tasks.append(task)

    def run(self):
        """运行调度器"""
        while self.tasks:
            task = self.tasks.popleft()
            try:
                # 执行到下一个 yield
                next(task)
                # 任务未完成,放回队列
                self.tasks.append(task)
            except StopIteration:
                # 任务完成
                pass

# 定义协程任务
def task1():
    """任务1:打印3次"""
    for i in range(3):
        print(f"任务1 - 第{i+1}次执行")
        yield  # 让出控制权

def task2():
    """任务2:打印5次"""
    for i in range(5):
        print(f"任务2 - 第{i+1}次执行")
        yield

def task3():
    """任务3:计算并打印"""
    total = 0
    for i in range(1, 6):
        total += i
        print(f"任务3 - 累计和: {total}")
        yield

# 运行调度器
print("=== 协程调度演示 ===")
scheduler = TaskScheduler()
scheduler.add_task(task1())
scheduler.add_task(task2())
scheduler.add_task(task3())
scheduler.run()

# 进阶:带延迟的任务调度
class TimedTaskScheduler:
    """支持延迟的任务调度器"""

    def __init__(self):
        self.tasks = []  # (next_run_time, task)

    def add_task(self, task, delay=0):
        """添加任务,可指定初始延迟"""
        next_run = time.time() + delay
        self.tasks.append((next_run, task))
        self.tasks.sort(key=lambda x: x[0])

    def run(self):
        """运行调度器"""
        while self.tasks:
            next_run, task = self.tasks.pop(0)

            # 等待到执行时间
            sleep_time = next_run - time.time()
            if sleep_time > 0:
                time.sleep(sleep_time)

            try:
                # 执行任务,获取下次延迟
                delay = next(task)
                # 重新调度
                next_run = time.time() + (delay if delay else 0.1)
                self.tasks.append((next_run, task))
                self.tasks.sort(key=lambda x: x[0])
            except StopIteration:
                pass

def periodic_task(name, count):
    """周期性任务"""
    for i in range(count):
        print(f"[{time.time():.2f}] {name} - 第{i+1}次")
        yield 0.5  # 下次执行延迟0.5秒

print("\n=== 定时调度演示 ===")
timed_scheduler = TimedTaskScheduler()
timed_scheduler.add_task(periodic_task("A", 3))
timed_scheduler.add_task(periodic_task("B", 4), delay=0.2)
# timed_scheduler.run()  # 取消注释以运行

4.3 生成器实现管道处理

# 使用生成器构建数据处理管道
def read_lines(filename):
    """读取文件行"""
    with open(filename, 'r', encoding='utf-8') as f:
        for line in f:
            yield line.strip()

def filter_comments(lines):
    """过滤注释行"""
    for line in lines:
        if not line.startswith('#'):
            yield line

def parse_csv(lines):
    """解析 CSV 格式"""
    for line in lines:
        if line:
            yield line.split(',')

def filter_columns(rows, min_columns=3):
    """过滤列数不足的行"""
    for row in rows:
        if len(row) >= min_columns:
            yield row

def convert_types(rows):
    """转换数据类型"""
    for row in rows:
        try:
            # 假设第一列是字符串,第二列是整数,第三列是浮点数
            converted = [
                row[0].strip(),
                int(row[1].strip()),
                float(row[2].strip())
            ]
            yield converted
        except (ValueError, IndexError):
            # 跳过转换失败的行
            continue

def calculate_stats(rows):
    """计算统计信息"""
    total = 0
    count = 0
    for row in rows:
        total += row[2]  # 第三列的数值
        count += 1
        yield row

    # 最后返回统计信息
    yield f"平均价格: {total/count:.2f}"

# 构建处理管道(创建测试数据)
test_data = """# 产品数据
苹果,10,3.5
香蕉,20,2.0
# 这是注释
橙子,15,4.0
无效行
葡萄,12,5.5
"""

with open('test_products.csv', 'w', encoding='utf-8') as f:
    f.write(test_data)

# 管道处理
print("=== 数据处理管道 ===")
pipeline = calculate_stats(
    convert_types(
        filter_columns(
            parse_csv(
                filter_comments(
                    read_lines('test_products.csv')
                )
            )
        )
    )
)

for result in pipeline:
    print(result)

# 更优雅的管道写法(使用生成器表达式)
def create_pipeline(filename):
    """使用生成器表达式构建管道"""
    lines = read_lines(filename)
    filtered = filter_comments(lines)
    parsed = parse_csv(filtered)
    columns_ok = filter_columns(parsed)
    converted = convert_types(columns_ok)
    return calculate_stats(converted)

# 管道模式的实际应用:日志分析
def analyze_logs(log_file):
    """分析日志文件,找出错误最多的IP"""
    from collections import Counter

    # 管道:读取 -> 过滤错误 -> 提取IP -> 统计
    error_lines = (line for line in open(log_file, encoding='utf-8') 
                   if 'ERROR' in line)
    ips = (line.split()[0] for line in error_lines)
    counter = Counter(ips)

    return counter.most_common(5)

# 创建测试日志
log_content = """192.168.1.1 - INFO - Request successful
192.168.1.2 - ERROR - Database connection failed
192.168.1.1 - ERROR - Timeout occurred
192.168.1.3 - INFO - Cache hit
192.168.1.2 - ERROR - Invalid request
192.168.1.1 - ERROR - Service unavailable
192.168.1.4 - ERROR - Authentication failed
"""

with open('test.log', 'w', encoding='utf-8') as f:
    f.write(log_content)

print("\n=== 日志分析结果 ===")
for ip, count in analyze_logs('test.log'):
    print(f"IP: {ip:15s} 错误次数: {count}")

五、迭代器工具库:itertools

5.1 无限迭代器

import itertools

print("=== 无限迭代器 ===")

# 1. count(start, step) - 无限计数
counter = itertools.count(10, 2)
print("count(10, 2):", [next(counter) for _ in range(5)])

# 2. cycle(iterable) - 无限循环
cycler = itertools.cycle(['A', 'B', 'C'])
print("cycle:", [next(cycler) for _ in range(7)])

# 3. repeat(object, times) - 重复对象
repeater = itertools.repeat('Hello', 3)
print("repeat:", list(repeater))

# 实战:生成ID
def generate_ids(prefix='USER'):
    """生成唯一ID"""
    for i in itertools.count(1):
        yield f"{prefix}_{i:06d}"

id_gen = generate_ids()
print(f"\n生成的ID: {[next(id_gen) for _ in range(5)]}")

5.2 组合迭代器

print("\n=== 组合迭代器 ===")

# 1. chain - 连接多个可迭代对象
result = itertools.chain([1, 2], ['a', 'b'], (True, False))
print(f"chain: {list(result)}")

# chain.from_iterable - 连接嵌套的可迭代对象
nested = [[1, 2], [3, 4, 5], [6]]
result = itertools.chain.from_iterable(nested)
print(f"chain.from_iterable: {list(result)}")

# 2. zip_longest - 按最长的迭代
list1 = [1, 2, 3]
list2 = ['a', 'b']
result = itertools.zip_longest(list1, list2, fillvalue=None)
print(f"zip_longest: {list(result)}")

# 3. product - 笛卡尔积
colors = ['red', 'blue']
sizes = ['S', 'M', 'L']
result = itertools.product(colors, sizes)
print(f"product: {list(result)}")

# 4. permutations - 排列
items = ['A', 'B', 'C']
print(f"permutations (2): {list(itertools.permutations(items, 2))}")

# 5. combinations - 组合
print(f"combinations (2): {list(itertools.combinations(items, 2))}")

# 6. combinations_with_replacement - 可重复组合
print(f"combinations_with_replacement (2): {list(itertools.combinations_with_replacement(items, 2))}")

5.3 过滤和分组

print("\n=== 过滤和分组迭代器 ===")

# 1. compress - 根据选择器过滤
data = ['A', 'B', 'C', 'D', 'E']
selectors = [True, False, True, False, True]
result = itertools.compress(data, selectors)
print(f"compress: {list(result)}")

# 2. dropwhile / takewhile - 条件过滤
numbers = [1, 3, 5, 2, 4, 6]
print(f"dropwhile (<4): {list(itertools.dropwhile(lambda x: x < 4, numbers))}")
print(f"takewhile (<4): {list(itertools.takewhile(lambda x: x < 4, numbers))}")

# 3. filterfalse - 过滤假值
mixed = [0, 1, False, True, '', 'hello', None, 42]
result = itertools.filterfalse(None, mixed)
print(f"filterfalse: {list(result)}")

# 4. groupby - 分组(需要先排序)
data = [('A', 1), ('A', 2), ('B', 3), ('B', 4), ('C', 5)]
# 按第一个元素分组
for key, group in itertools.groupby(data, key=lambda x: x[0]):
    print(f"  groupby '{key}': {list(group)}")

# 实战:数据分组处理
def group_and_aggregate():
    """按部门统计工资"""
    employees = [
        ('Alice', 'Sales', 5000),
        ('Bob', 'Sales', 5500),
        ('Charlie', 'IT', 7000),
        ('David', 'IT', 7500),
        ('Eve', 'HR', 4500)
    ]

    # 先按部门排序(groupby 要求)
    employees.sort(key=lambda x: x[1])

    print("\n部门工资统计:")
    for dept, group in itertools.groupby(employees, key=lambda x: x[1]):
        group_list = list(group)
        total_salary = sum(emp[2] for emp in group_list)
        avg_salary = total_salary / len(group_list)
        print(f"  {dept}: 人数={len(group_list)}, 总工资={total_salary}, 平均={avg_salary:.0f}")

group_and_aggregate()

5.4 实战:itertools 解决复杂问题

print("\n=== itertools 实战应用 ===")

# 1. 滑动窗口
def sliding_window(iterable, n):
    """返回大小为 n 的滑动窗口"""
    iterators = itertools.tee(iterable, n)
    for i, it in enumerate(iterators):
        # 每个迭代器跳过不同数量的元素
        for _ in range(i):
            next(it, None)
    return zip(*iterators)

data = [1, 2, 3, 4, 5, 6]
print(f"滑动窗口(3): {list(sliding_window(data, 3))}")

# 2. 计算移动平均
def moving_average(iterable, n):
    """使用滑动窗口计算移动平均"""
    for window in sliding_window(iterable, n):
        yield sum(window) / n

prices = [10, 12, 15, 14, 18, 20, 17]
print(f"移动平均(3): {list(moving_average(prices, 3))}")

# 3. 生成所有可能的密码组合
def generate_passwords(charset, max_length):
    """生成指定字符集的所有可能密码"""
    for length in range(1, max_length + 1):
        for combo in itertools.product(charset, repeat=length):
            yield ''.join(combo)

# 注意:这会生成很多组合,谨慎使用
chars = 'abc'
print(f"密码组合 (a,b,c, 长度1-2): {list(generate_passwords(chars, 2))}")

# 4. 配对处理(前后对比)
def pairwise_comparison(iterable):
    """比较相邻元素"""
    a, b = itertools.tee(iterable)
    next(b, None)
    for prev, curr in zip(a, b):
        yield (prev, curr, curr - prev if isinstance(curr, (int, float)) else None)

values = [10, 15, 13, 20, 18]
print("\n相邻元素比较:")
for prev, curr, diff in pairwise_comparison(values):
    print(f"  {prev} -> {curr}: 变化 {diff:+d}")

# 5. 分块处理大列表
def chunked(iterable, chunk_size):
    """将可迭代对象分成固定大小的块"""
    iterator = iter(iterable)
    return iter(lambda: list(itertools.islice(iterator, chunk_size)), [])

large_list = list(range(20))
print(f"\n分块处理 (大小5):")
for chunk in chunked(large_list, 5):
    print(f"  处理块: {chunk}")

六、性能与内存优化

6.1 生成器 vs 列表的性能对比

import sys
import timeit

def memory_comparison():
    """比较生成器和列表的内存占用"""
    n = 1000000

    # 列表
    list_squares = [x * x for x in range(n)]
    list_size = sys.getsizeof(list_squares)
    # 还要计算元素占用的内存
    element_size = sum(sys.getsizeof(x) for x in list_squares[:100]) / 100 * n
    total_list_memory = list_size + element_size

    # 生成器
    gen_squares = (x * x for x in range(n))
    gen_size = sys.getsizeof(gen_squares)

    print(f"列表推导式内存占用: {total_list_memory / 1024 / 1024:.2f} MB")
    print(f"生成器表达式内存占用: {gen_size} 字节")
    print(f"内存节省: {(1 - gen_size/total_list_memory) * 100:.2f}%")

def speed_comparison():
    """比较生成器和列表的迭代速度"""
    n = 1000000

    # 创建但不迭代
    list_time = timeit.timeit(
        lambda: [x * x for x in range(n)],
        number=10
    )

    gen_time = timeit.timeit(
        lambda: (x * x for x in range(n)),
        number=10
    )

    print(f"\n创建时间对比(10次):")
    print(f"  列表推导式: {list_time:.4f}s")
    print(f"  生成器表达式: {gen_time:.4f}s")
    print(f"  生成器更快: {list_time/gen_time:.2f}倍")

    # 实际迭代使用
    def iterate_list():
        return sum([x * x for x in range(n)])

    def iterate_gen():
        return sum(x * x for x in range(n))

    list_iter_time = timeit.timeit(iterate_list, number=10)
    gen_iter_time = timeit.timeit(iterate_gen, number=10)

    print(f"\n迭代计算时间(10次):")
    print(f"  列表方式: {list_iter_time:.4f}s")
    print(f"  生成器方式: {gen_iter_time:.4f}s")

# 运行对比
memory_comparison()
speed_comparison()

6.2 惰性求值的威力

# 演示惰性求值如何节省计算
def expensive_operation(x):
    """模拟耗时操作"""
    import time
    time.sleep(0.001)  # 模拟1ms的计算
    return x * x

print("\n=== 惰性求值示例 ===")

# 急需求值:立即计算所有结果
def eager_approach(n):
    results = [expensive_operation(i) for i in range(n)]
    # 只使用前5个结果
    return sum(results[:5])

# 惰性求值:只计算需要的部分
def lazy_approach(n):
    results = (expensive_operation(i) for i in range(n))
    # 同样只使用前5个结果
    return sum(itertools.islice(results, 5))

n = 100
eager_time = timeit.timeit(lambda: eager_approach(n), number=1)
lazy_time = timeit.timeit(lambda: lazy_approach(n), number=1)

print(f"急需求值时间: {eager_time:.4f}s (计算了 {n} 次)")
print(f"惰性求值时间: {lazy_time:.4f}s (只计算了 5 次)")
print(f"性能提升: {eager_time/lazy_time:.2f}倍")

# 无限序列的处理
def process_infinite_sequence():
    """处理无限序列,只取需要的部分"""
    # 生成所有偶数
    even_numbers = (x for x in itertools.count() if x % 2 == 0)

    # 取前10个
    first_ten = list(itertools.islice(even_numbers, 10))
    print(f"\n前10个偶数: {first_ten}")

    # 找第一个大于1000的偶数
    even_numbers = (x for x in itertools.count() if x % 2 == 0)
    first_over_1000 = next(x for x in even_numbers if x > 1000)
    print(f"第一个大于1000的偶数: {first_over_1000}")

process_infinite_sequence()

七、最佳实践与常见陷阱

7.1 常见陷阱和注意事项

print("=== 常见陷阱 ===")

# 陷阱1:生成器只能迭代一次
def trap1():
    gen = (x for x in range(3))
    print(f"第一次迭代: {list(gen)}")
    print(f"第二次迭代: {list(gen)}")  # 空列表!

# 陷阱2:在迭代时修改集合
def trap2():
    numbers = [1, 2, 3, 4, 5]
    # 错误:在迭代时删除元素
    try:
        for num in numbers:
            if num % 2 == 0:
                numbers.remove(num)  # 会导致跳过元素
    except Exception as e:
        print(f"错误: {e}")

    # 正确:创建新列表或使用列表推导式
    numbers = [1, 2, 3, 4, 5]
    numbers = [num for num in numbers if num % 2 != 0]
    print(f"过滤后: {numbers}")

# 陷阱3:yield 在 try-finally 中
def trap3():
    def generator_with_finally():
        try:
            yield 1
            yield 2
        finally:
            print("finally 块执行")

    gen = generator_with_finally()
    print(f"获取: {next(gen)}")
    print(f"获取: {next(gen)}")
    # 如果生成器没有完全消费,finally 不会执行!
    print("生成器未完全消费,finally 未执行")

    # 正确:确保生成器被完全消费或显式关闭
    gen = generator_with_finally()
    next(gen)
    gen.close()  # 显式关闭会触发 finally

# 陷阱4:生成器中的变量作用域
def trap4():
    funcs = []
    for i in range(3):
        # 错误:lambda 捕获的是变量 i 的引用
        funcs.append(lambda: i)
    print(f"错误结果: {[f() for f in funcs]}")  # [2, 2, 2]

    # 正确:使用默认参数固定值
    funcs = []
    for i in range(3):
        funcs.append(lambda x=i: x)
    print(f"正确结果: {[f() for f in funcs]}")  # [0, 1, 2]

trap1()
trap2()
trap3()
trap4()

7.2 最佳实践指南

"""
迭代器与生成器最佳实践:

1. ✓ 使用生成器处理大数据
   - 避免一次性加载所有数据到内存
   - 使用生成器表达式代替列表推导式

2. ✓ 善用 itertools
   - 复杂的迭代逻辑优先使用 itertools
   - 避免重复造轮子

3. ✓ 生成器函数设计
   - 保持生成器函数简单,单一职责
   - 使用 yield from 简化委托

4. ✓ 性能考虑
   - 小数据量用列表(访问快)
   - 大数据量用生成器(内存省)
   - 需要多次迭代时转换为列表

5. ✓ 错误处理
   - 使用 try-finally 确保资源释放
   - 显式 close() 未完全消费的生成器

6. ✓ 代码可读性
   - 复杂生成器拆分为多个小生成器
   - 使用管道模式组织数据流
"""

# 示例:符合最佳实践的数据处理类
class DataPipeline:
    """符合最佳实践的数据处理管道"""

    def __init__(self, source):
        self.source = source
        self._pipeline = None

    def filter(self, predicate):
        """添加过滤器"""
        if self._pipeline is None:
            self._pipeline = (x for x in self.source if predicate(x))
        else:
            self._pipeline = (x for x in self._pipeline if predicate(x))
        return self

    def map(self, transform):
        """添加转换器"""
        if self._pipeline is None:
            self._pipeline = (transform(x) for x in self.source)
        else:
            self._pipeline = (transform(x) for x in self._pipeline)
        return self

    def take(self, n):
        """取前 n 个元素"""
        if self._pipeline is None:
            self._pipeline = itertools.islice(self.source, n)
        else:
            self._pipeline = itertools.islice(self._pipeline, n)
        return self

    def execute(self):
        """执行管道"""
        if self._pipeline is None:
            return iter(self.source)
        return self._pipeline

    def to_list(self):
        """转换为列表(谨慎使用,会消耗生成器)"""
        return list(self.execute())

# 使用示例
print("\n=== 数据处理管道示例 ===")
data = range(100)
result = (DataPipeline(data)
          .filter(lambda x: x % 2 == 0)   # 偶数
          .map(lambda x: x * x)            # 平方
          .filter(lambda x: x > 1000)      # 大于1000
          .take(5)                         # 前5个
          .to_list())

print(f"处理结果: {result}")

八、总结

迭代器和生成器是 Python 最强大、最优雅的特性之一。它们体现了 Python "惰性求值" 和 "流式处理" 的哲学。

核心要点: - 迭代器协议:__iter__() 和 __next__() 是迭代的基础 - 生成器函数:使用 yield 创建,自动实现迭代器协议 - 生成器表达式:内存友好的列表推导式替代品 - yield from:简化生成器委托和递归 - itertools:迭代器操作的瑞士军刀

选择建议: - 处理大数据流 → 生成器 - 需要多次访问数据 → 列表 - 复杂迭代逻辑 → itertools - 异步/协程 → 生成器协程(Python 3.5+ 可用 async/await)

掌握迭代器和生成器,你就能写出更优雅、更高效的 Python 代码。它们是 Python 进阶之路上的重要里程碑。


本文由 尚先生 原创,转载请注明出处。

📖相关推荐

评论

0
暂无评论,来发表第一条评论吧

发表评论

登录 后发表评论