You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

迭代器仅支持主线程时,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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.21 02:23:14