Python asyncio:Queue.join() 仅在未引发异常时完成,为什么?(上下文:编写异步映射函数)

Jie*_*ong 6 python python-asyncio

我一直在尝试编写一个异步版本map我一直在尝试用 Python

为此,我使用带有生产者/消费者的队列。

起初它似乎运行良好,但无一例外。

特别是,如果我使用queue.join(),它在没有异常时运行良好,但在异常情况下会阻塞。如果我使用gather(*tasks),它在出现异常时效果很好,但如果没有则阻塞。

所以有时它只会完成,我只是不明白为什么。

这是我实现的代码:

import asyncio
from asyncio import Queue
from typing import Iterable, Callable, TypeVar

Input = TypeVar("Input")
Output = TypeVar("Output")
STOP = object()


def parallel_map(func: Callable[[Input], Output], iterable: Iterable[Input]) -> Iterable[Output]:
    """
    Parallel version of `map`, backed by asyncio.
    Only suitable to do IO in parallel (not for CPU intensive tasks, otherwise it will block).
    """

    number_of_parallel_calls = 9

    async def worker(input_queue: Queue, output_queue: Queue) -> None:
        while True:
            data = await input_queue.get()
            try:
                output = func(data)
                # Simulate an exception:
                # raise RuntimeError("")
                output_queue.put_nowait(output)
            finally:
                input_queue.task_done()

    async def group_results(output_queue: Queue) -> Iterable[Output]:
        output = []
        while True:
            item = await output_queue.get()
            if item is not STOP:
                output.append(item)
            output_queue.task_done()
            if item is STOP:
                break
        return output

    async def procedure() -> Iterable[Output]:
        # First, produce a queue of inputs
        input_queue: Queue = asyncio.Queue()
        for i in iterable:
            input_queue.put_nowait(i)

        # Then, assign a pool of tasks to consume it (and also produce outputs in a new queue)
        output_queue: Queue = asyncio.Queue()
        tasks = []
        for _ in range(number_of_parallel_calls):
            task = asyncio.create_task(worker(input_queue, output_queue))
            tasks.append(task)

        # Wait for the input queue to be fully consumed (only works if no exception occurs in the tasks), blocks otherwise.
        await input_queue.join()
        # Gather tasks, only works when an exception is raised in a task, blocks otherwise
        # asyncio.gather(*tasks)

        for task in tasks:
            task.cancel()

        # Indicate that the output queue is complete, to stop the worker
        output_queue.put_nowait(STOP)

        # Consume the output_queue, and return its data as a list
        group_results_task = asyncio.create_task(group_results(output_queue))
        await output_queue.join()
        output = await group_results_task
        return output

    return asyncio.run(procedure())

if __name__ == "__main__()":
    def my_function(x):
        return x * x

    data = [1, 2, 3, 4]
    print(parallel_map(my_function, data))
Run Code Online (Sandbox Code Playgroud)

我认为我误解了Python asyncio 的基本但重要的部分,但不确定是什么。

jup*_*bjy 7

问题是,你没有捕获异常。

来自Python 文档

每当将项目添加到队列中时,未完成任务的计数就会增加。每当消费者协程调用task_done()以指示该项目已被检索并且其上的所有工作都已完成时,计数就会减少。当未完成任务的计数降至零时, join() 会解除阻塞。

因此Queue,本质上是计算 上的调用次数put(),并在每次调用时将计数器减 1 task_done()。如果工作进程在处理所有队列之前停止,您将被阻塞Queue.join()。

在您的工人代码处:

    async def worker(input_queue: Queue, output_queue: Queue) -> None:
        while True:
            data = await input_queue.get()
            try:
                output = func(data)
                output_queue.put_nowait(output)
            finally:
                input_queue.task_done()
Run Code Online (Sandbox Code Playgroud)

你的工作线程在遇到Exception时会停止,因为try-finally 只保证清理,而不是实际的错误处理。

因此,您的情况发生的是:

  1. 每次Queue.put()调用都会增加内部计数器,假设我们调用了n多次。
  2. 工人们开始呼吁Queue.task_done()减少内部计数器。
  3. 当遇到错误时,worker 将在执行完 block 后停止finally。
  4. 现在,如果所有工作人员停止,Queue.task_done()调用计数n'为n' < n,内部计数器仍为正值。
  5. Queue.join()无限期挂起,直到内部计数器为 0,这种情况永远不会发生,因为所有工作人员都死了。

这是一个设计缺陷。


其他有用的更改

为了易于实施、减少故障和提高性能,需要更改多个设计因素。

请注意,这是我使用 python 的经验,所以不要将此视为具体事实。

对于设计因素,我做了以下更改:

  • 如果只是运行的话是没有意义的function,所以最好coroutine也支持一下。
  • input_queue保证在 之前被填充worker。检查Queue.empty()足以确定循环结束。
  • 如果你从填充开始input_queue,那么不需要哨兵,你知道给出了多长时间iterable,由queue.qszie()。
  • 使用await Queue.put()而不是put_nowait(),您无法确定Queue在您放置它的精确时间是否不可用。
  • 使用额外的索引参数保持输入和输出的顺序,您也无法确定并发任务的输出顺序。
  • 不是在 中实现故障安全Exception,而是将错误放入结果中并处理所有队列,然后只需根据用户的选择重新引发它。
  • 在任务生成中使用genexpfor是该任务不需要的 - 并且list.append会影响脚本的性能。
  • 不要Queue从导入,如果它来自或任何其他具有内置对象的库,asyncio它不会向用户发出足够的警告。queueQueue
  • 完全不需要等待queue.join运行。group_resultsawait

功能代码:

import asyncio


def parallel_map(func, iterable, concurrent_limit=2, raise_error=False):
    async def worker(input_queue: asyncio.Queue, output_queue: asyncio.Queue):

        while not input_queue.empty():
            idx, item = await input_queue.get()
            try:
                # Support both coroutine and function. Coroutine function I mean!
                if asyncio.iscoroutinefunction(func):
                    output = await func(item)
                else:
                    output = func(item)
                await output_queue.put((idx, output))

            except Exception as err:
                await output_queue.put((idx, err))

            finally:
                input_queue.task_done()

    async def group_results(input_size, output_queue: asyncio.Queue):
        output = {}  # using dict to remove the need to sort list

        for _ in range(input_size):
            idx, val = await output_queue.get()  # gets tuple(idx, result)
            output[idx] = val
            output_queue.task_done()

        return [output[i] for i in range(input_size)]

    async def procedure():
        # populating input queue
        input_queue: asyncio.Queue = asyncio.Queue()
        for idx, item in enumerate(iterable):
            input_queue.put_nowait((idx, item))
        
        # Remember size before using Queue
        input_size = input_queue.qsize()

        # Generate task pool, and start collecting data.
        output_queue: asyncio.Queue = asyncio.Queue()
        result_task = asyncio.create_task(group_results(input_size, output_queue))
        tasks = [
            asyncio.create_task(worker(input_queue, output_queue))
            for _ in range(concurrent_limit)
        ]
        
        # Wait for tasks complete
        await asyncio.gather(*tasks)
        
        # Wait for result fetching
        results = await result_task
        
        # Re-raise errors at once if raise_error
        if raise_error and (errors := [err for err in results if isinstance(err, Exception)]):
            # noinspection PyUnboundLocalVariable
            raise Exception(errors)  # It never runs before assignment, safe to ignore.

        return results

    return asyncio.run(procedure())
Run Code Online (Sandbox Code Playgroud)

测试代码:

if __name__ == "__main__":
    import random
    import time

    data = [1, 2, 3]
    err_data = [1, 'yo', 3]

    def test_normal_function(data_, raise_=False):
        def my_function(x):
            t = random.uniform(1, 2)
            print(f"Sleep {t:.3} start")

            time.sleep(t)
            print(f"Awake after {t:.3}")

            return x * x

        print(f"Normal function: {parallel_map(my_function, data_, raise_error=raise_)}\n")

    def test_coroutine(data_, raise_=False):
        async def my_coro(x):
            t = random.uniform(1, 2)
            print(f"Coroutine sleep {t:.3} start")

            await asyncio.sleep(t)
            print(f"Coroutine awake after {t:.3}")

            return x * x

        print(f"Coroutine {parallel_map(my_coro, data_, raise_error=raise_)}\n")

    # Test starts
    print(f"Test for data {data}:")
    test_normal_function(data)
    test_coroutine(data)

    print(f"Test for data {err_data} without raise:")
    test_normal_function(err_data)
    test_coroutine(err_data)

    print(f"Test for data {err_data} with raise:")
    test_normal_function(err_data, True)
    test_coroutine(err_data, True)  # this line will not run, but works same.
Run Code Online (Sandbox Code Playgroud)

上面将测试 和 的以下function条件coroutine:

  • 正常数据
  • 数据有缺陷,不引发错误
  • 有缺陷的数据,提出遇到的所有错误

即使出现异常,这也不会取消任务,而是处理所有队列。

输出:

    async def worker(input_queue: Queue, output_queue: Queue) -> None:
        while True:
            data = await input_queue.get()
            try:
                output = func(data)
                output_queue.put_nowait(output)
            finally:
                input_queue.task_done()
Run Code Online (Sandbox Code Playgroud)

请注意,我设置了concurrent_limit2 来演示协程等待可用的工作线程。这就是为什么3 个协程任务中只有一个没有立即运行。

从输出中您还可以看到一些任务先于其他任务完成,但结果是按顺序排列的。


聚苯乙烯

如果您由于通过类型提示违反 PEP-8 行限制而单独导入Queue,则可以添加类型提示,如下所示:

async def worker(input_queue, output_queue) -> None:
    input_queue: asyncio.Queue
    output_queue: asyncio.Queue
Run Code Online (Sandbox Code Playgroud)

或者

async def worker(
        input_queue: asyncio.Queue,
        output_queue: asyncio.Queue
    ) -> None:
Run Code Online (Sandbox Code Playgroud)

虽然它不如原始方式那么干净,但这将有助于其他人阅读您的代码。

  • @Jiehong它确实通知队列当前任务已完成,但随后它_停止处理任务_。因此,生产者排队的其他任务不会被工作线程拾取,也不会调用“task_done()”。由于您的测试代码向所有工作人员添加了异常引发,因此没有人留下来处理这些项目,这会阻止“join()”完成。 (2认同)