1. 为什么需要自定义迭代器
在Python中,迭代器(Iterator)是一种设计模式,它允许我们按顺序访问集合中的元素而不需要暴露其底层表示。内置的列表、元组、字典等数据类型都实现了迭代器协议,可以直接用于for循环。但当我们处理自定义数据结构或需要特殊遍历逻辑时,内置迭代器往往无法满足需求。
举个例子,假设我们正在开发一个树形结构的数据存储系统,标准的深度优先或广度优先遍历可能都不符合业务需求。或者我们有一个大型数据集需要分块处理,每次迭代返回固定大小的数据块而非单个元素。这些场景下,自定义迭代器就变得非常必要。
提示:迭代器模式的核心价值在于将集合的遍历逻辑与业务逻辑解耦,使代码更清晰且易于维护。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 迭代器协议详解
2.1 Python中的迭代器协议
Python的迭代器协议由两个特殊方法组成:
__iter__(): 返回迭代器对象本身__next__(): 返回容器中的下一个元素,如果没有更多元素则抛出StopIteration异常
任何实现了这两个方法的对象都可以被视为迭代器。下面是一个最简单的实现示例:
python复制class SimpleIterator:
def __init__(self, max_num):
self.max_num = max_num
self.current = 0
def __iter__(self):
return self
def __next__(self):
if self.current < self.max_num:
num = self.current
self.current += 1
return num
raise StopIteration
2.2 可迭代对象与迭代器的区别
初学者常常混淆这两个概念:
- 可迭代对象(Iterable):实现了
__iter__()方法的对象,可以返回一个迭代器 - 迭代器(Iterator):实现了
__iter__()和__next__()方法的对象
列表是可迭代对象但不是迭代器,因为它没有实现__next__()方法。调用iter()函数时,列表会返回一个专门的列表迭代器对象。
3. 自定义迭代器实战案例
3.1 分块迭代器实现
处理大型数据集时,我们经常需要分块读取数据。下面实现一个分块迭代器:
python复制class ChunkIterator:
def __init__(self, data, chunk_size):
self.data = data
self.chunk_size = chunk_size
self.index = 0
def __iter__(self):
return self
def __next__(self):
if self.index >= len(self.data):
raise StopIteration
chunk = self.data[self.index : self.index + self.chunk_size]
self.index += self.chunk_size
return chunk
使用示例:
python复制data = list(range(100))
for chunk in ChunkIterator(data, 10):
print(chunk) # 每次输出10个元素的列表
3.2 树形结构迭代器
对于树形结构,我们可以实现多种遍历方式的迭代器。以下是深度优先遍历的实现:
python复制class TreeNode:
def __init__(self, value):
self.value = value
self.children = []
def add_child(self, node):
self.children.append(node)
class DepthFirstIterator:
def __init__(self, root):
self.stack = [root]
def __iter__(self):
return self
def __next__(self):
if not self.stack:
raise StopIteration
node = self.stack.pop()
self.stack.extend(reversed(node.children))
return node.value
4. 高级迭代器技巧
4.1 使用生成器简化迭代器实现
Python的生成器函数(使用yield关键字)会自动实现迭代器协议,可以大大简化代码:
python复制def chunk_generator(data, chunk_size):
for i in range(0, len(data), chunk_size):
yield data[i:i + chunk_size]
4.2 带状态的迭代器
有时我们需要迭代器记住一些状态信息。例如,一个可以重置的迭代器:
python复制class ResettableIterator:
def __init__(self, data):
self.data = data
self.reset()
def reset(self):
self.index = 0
def __iter__(self):
return self
def __next__(self):
if self.index >= len(self.data):
raise StopIteration
item = self.data[self.index]
self.index += 1
return item
4.3 迭代器组合模式
我们可以创建组合多个迭代器的超级迭代器:
python复制class CombinedIterator:
def __init__(self, *iterables):
self.iterables = iterables
self.current_iter = iter(iterables[0])
self.iter_index = 0
def __iter__(self):
return self
def __next__(self):
try:
return next(self.current_iter)
except StopIteration:
self.iter_index += 1
if self.iter_index >= len(self.iterables):
raise
self.current_iter = iter(self.iterables[self.iter_index])
return next(self)
5. 性能优化与注意事项
5.1 惰性求值的优势
迭代器最大的优势是惰性求值(Lazy Evaluation),只在需要时生成数据,这对处理大型数据集或无限序列特别有用:
python复制def infinite_sequence():
num = 0
while True:
yield num
num += 1
5.2 内存效率对比
比较列表和迭代器的内存使用:
python复制import sys
# 列表一次性存储所有元素
big_list = [x for x in range(1000000)]
print(sys.getsizeof(big_list)) # 约9MB
# 迭代器只存储当前状态
iterator = (x for x in range(1000000))
print(sys.getsizeof(iterator)) # 仅128字节
5.3 常见陷阱与解决方案
-
迭代器耗尽问题:迭代器只能遍历一次,再次迭代需要重新创建
python复制it = iter([1, 2, 3]) list(it) # [1, 2, 3] list(it) # [] 已经耗尽 -
修改迭代中的集合:在迭代过程中修改集合会导致未定义行为
python复制# 错误示例 lst = [1, 2, 3] for item in lst: lst.remove(item) # 可能导致跳过元素或报错 -
无限迭代器:没有终止条件的迭代器会导致无限循环
python复制# 安全做法:使用itertools.islice限制数量 from itertools import islice for num in islice(infinite_sequence(), 100): print(num)
6. 实际应用场景分析
6.1 数据库查询结果分页
自定义迭代器非常适合处理数据库查询结果的分页获取:
python复制class DatabasePaginator:
def __init__(self, query, page_size=100):
self.query = query
self.page_size = page_size
self.offset = 0
def __iter__(self):
return self
def __next__(self):
results = execute_query(self.query, self.offset, self.page_size)
if not results:
raise StopIteration
self.offset += self.page_size
return results
6.2 日志文件处理
处理大型日志文件时,可以逐行或按时间窗口迭代:
python复制class LogFileIterator:
def __init__(self, filename, time_window=60):
self.filename = filename
self.time_window = time_window # 分钟
def __iter__(self):
with open(self.filename) as f:
current_window = None
window_lines = []
for line in f:
timestamp = parse_timestamp(line)
if current_window is None:
current_window = timestamp
if (timestamp - current_window).total_seconds() > self.time_window * 60:
yield window_lines
window_lines = []
current_window = timestamp
window_lines.append(line)
if window_lines:
yield window_lines
6.3 机器学习数据流处理
在机器学习中,自定义迭代器可用于实现复杂的数据流:
python复制class DataStream:
def __init__(self, source, preprocessors=None, augmentations=None):
self.source = source
self.preprocessors = preprocessors or []
self.augmentations = augmentations or []
def __iter__(self):
for item in self.source:
# 应用预处理
for preprocessor in self.preprocessors:
item = preprocessor(item)
# 应用数据增强
if self.augmentations:
yield item
for augment in self.augmentations:
yield augment(item)
else:
yield item
