test/app/core/executor.py

186 lines
5.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

# -*- coding: utf-8 -*-
"""
异步执行器模块
用于将同步的模型推理调用放入线程池执行,避免阻塞事件循环,
实现真正的多路并发处理。
设计要点:
1. 使用 ThreadPoolExecutor 而非 ProcessPoolExecutor
- 模型已加载在内存中,进程间无法共享
- GPU操作会自动释放GIL,线程池足以实现并发
2. 对于流式生成器,使用 asyncio.Queue 实现异步迭代
3. 线程池大小根据使用场景配置:
- CPU推理:受GIL限制,多线程并发收益有限,但可以让I/O不阻塞
- GPU推理:CUDA操作释放GIL,可以实现真正并发
"""
import os
import asyncio
import logging
from concurrent.futures import ThreadPoolExecutor
from typing import Callable, TypeVar, Generator, AsyncGenerator, Optional
from functools import partial
logger = logging.getLogger(__name__)
# 类型变量
T = TypeVar("T")
# 全局线程池执行器
# 默认线程数:max(4, CPU核心数),可通过环境变量覆盖
_DEFAULT_WORKERS = max(4, os.cpu_count() or 4)
_MAX_WORKERS = int(os.getenv("INFERENCE_THREAD_POOL_SIZE", str(_DEFAULT_WORKERS)))
_executor: Optional[ThreadPoolExecutor] = None
def get_executor() -> ThreadPoolExecutor:
"""获取全局线程池执行器(懒加载)"""
global _executor
if _executor is None:
_executor = ThreadPoolExecutor(
max_workers=_MAX_WORKERS,
thread_name_prefix="inference_worker"
)
logger.info(f"推理线程池已创建,最大工作线程数: {_MAX_WORKERS}")
return _executor
def shutdown_executor():
"""关闭线程池执行器"""
global _executor
if _executor is not None:
_executor.shutdown(wait=True)
_executor = None
logger.info("推理线程池已关闭")
async def run_sync(func: Callable[..., T], *args, **kwargs) -> T:
"""
在线程池中执行同步函数,不阻塞事件循环
Args:
func: 同步函数
*args: 位置参数
**kwargs: 关键字参数
Returns:
函数返回值
Example:
result = await run_sync(model.generate, input=audio_array, cache=cache)
"""
loop = asyncio.get_running_loop()
executor = get_executor()
# 使用 partial 绑定参数
if kwargs:
func_with_args = partial(func, *args, **kwargs)
else:
func_with_args = partial(func, *args) if args else func
return await loop.run_in_executor(executor, func_with_args)
async def run_sync_generator(
generator_func: Callable[..., Generator[T, None, None]],
*args,
**kwargs
) -> AsyncGenerator[T, None]:
"""
将同步生成器转换为异步生成器,在线程池中执行
用于流式处理等需要逐步产出结果的场景。
Args:
generator_func: 返回生成器的同步函数
*args: 位置参数
**kwargs: 关键字参数
Yields:
生成器的每个产出值
Example:
async for chunk in run_sync_generator(model.inference_sft, text, voice, stream=True):
await websocket.send_bytes(chunk)
"""
loop = asyncio.get_running_loop()
executor = get_executor()
queue: asyncio.Queue = asyncio.Queue()
# 标记生成器结束的哨兵值
_SENTINEL = object()
def producer():
"""在线程中运行生成器,将结果放入队列"""
try:
gen = generator_func(*args, **kwargs)
for item in gen:
# 使用 call_soon_threadsafe 安全地将结果放入队列
loop.call_soon_threadsafe(queue.put_nowait, item)
except BaseException as e:
# 发生异常时,记录日志并将异常放入队列
logger.error(f"生成器执行异常: {type(e).__name__}: {e}")
loop.call_soon_threadsafe(queue.put_nowait, e)
finally:
# 发送结束标记
loop.call_soon_threadsafe(queue.put_nowait, _SENTINEL)
# 在线程池中启动生产者
future = executor.submit(producer)
try:
while True:
item = await queue.get()
if item is _SENTINEL:
break
# 使用 BaseException 捕获所有异常类型(包括 KeyboardInterrupt 等)
if isinstance(item, BaseException):
raise item
yield item
finally:
# 确保检查线程是否有未捕获的异常
if future.done():
try:
# 如果线程已完成,检查是否有异常
future.result()
except Exception as e:
logger.error(f"生成器线程异常: {type(e).__name__}: {e}")
elif not future.cancelled():
future.cancel()
class AsyncInferenceWrapper:
"""
异步推理包装器
将同步的模型推理方法包装为异步方法,方便复用。
Example:
wrapper = AsyncInferenceWrapper(asr_engine.realtime_model)
result = await wrapper.generate(input=audio_array, cache=cache)
"""
def __init__(self, model):
self._model = model
async def generate(self, *args, **kwargs):
"""异步调用模型的 generate 方法"""
return await run_sync(self._model.generate, *args, **kwargs)
async def inference_sft(self, *args, **kwargs):
"""异步流式调用模型的 inference_sft 方法"""
async for item in run_sync_generator(self._model.inference_sft, *args, **kwargs):
yield item
async def inference_zero_shot(self, *args, **kwargs):
"""异步流式调用模型的 inference_zero_shot 方法"""
async for item in run_sync_generator(self._model.inference_zero_shot, *args, **kwargs):
yield item