迭代器仅支持主线程时,Python多进程Pool.imap的替代实现方案
问题解决:主线程专属迭代器的多进程任务分发
问题背景
现有一个仅能在主线程工作的迭代器GraphIterator(其他线程/进程调用会抛出断言异常),需要将迭代器生成的每个Graph对象对应的计算任务get_number_from_graph分发到多进程执行(单任务计算成本远高于迭代成本)。要求不得修改Graph、GraphIterator、get_number_from_graph,仅使用Python标准库修改代码,解决运行报错并得到预期输出。
原示例代码
from multiprocessing import Pool import threading class Graph: def __init__(self, num_vertices): self._num_vertices = num_vertices class GraphIterator: def __init__(self, num_graphs): self._num_graphs = num_graphs self._current_graph = 0 def __iter__(self): return self def __next__(self): assert threading.current_thread() is threading.main_thread(), 'iterator only works on the main thread' if self._current_graph < self._num_graphs: self._current_graph += 1 return Graph(self._current_graph) else: raise StopIteration def get_number_from_graph(graph): return graph._num_vertices if __name__ == '__main__': num_graphs = 100 print('Sequential result:', sum(get_number_from_graph(g) for g in GraphIterator(num_graphs))) print('Parallel result: ', end='') result = 0 with Pool(processes=None) as pool: for t in pool.imap(get_number_from_graph, GraphIterator(num_graphs)): result += t print(result)
当前运行错误输出
Sequential result: 5050 Parallel result: Traceback (most recent call last): File "/home/rburing/src/gcaops/multiprocessing_issue.py", line 33, in <module> for t in pool.imap(get_number_from_graph, GraphIterator(num_graphs)): File "/usr/lib/python3.10/multiprocessing/pool.py", line 873, in next raise value AssertionError: iterator only works on the main thread
解决方案思路
报错原因是pool.imap会在子进程中尝试拉取迭代器的元素,触发了GraphIterator的主线程断言。解决核心是确保迭代器的遍历完全在主线程完成,将生成的所有Graph对象先收集到列表中,再将列表传给多进程池处理。
修改后的代码
from multiprocessing import Pool import threading class Graph: def __init__(self, num_vertices): self._num_vertices = num_vertices class GraphIterator: def __init__(self, num_graphs): self._num_graphs = num_graphs self._current_graph = 0 def __iter__(self): return self def __next__(self): assert threading.current_thread() is threading.main_thread(), 'iterator only works on the main thread' if self._current_graph < self._num_graphs: self._current_graph += 1 return Graph(self._current_graph) else: raise StopIteration def get_number_from_graph(graph): return graph._num_vertices if __name__ == '__main__': num_graphs = 100 print('Sequential result:', sum(get_number_from_graph(g) for g in GraphIterator(num_graphs))) print('Parallel result: ', end='') result = 0 # 先在主线程遍历迭代器,收集所有Graph对象 graphs = list(GraphIterator(num_graphs)) with Pool(processes=None) as pool: # 使用map处理已收集的列表,任务分发到多进程 for t in pool.map(get_number_from_graph, graphs): result += t print(result)
预期输出
Sequential result: 5050 Parallel result: 5050
内容的提问来源于stack exchange,提问作者Ricardo Buring
相关产品推荐
相关产品推荐

