# -*- coding: utf-8 -*- """Runtime router for pooled ASR execution.""" from __future__ import annotations import asyncio import threading from dataclasses import dataclass from enum import Enum from typing import Awaitable, Callable, Optional import torch from app.core.accelerator import get_accelerator_info from app.core.config import settings from app.core.device import detect_device from app.core.executor import run_sync from app.services.asr.engines import ASRFullResult, BaseASREngine from app.services.asr.manager import get_model_manager from app.services.asr.qwenasr_rust import is_qwenasr_rust_available from .local_pool import LocalEnginePool class RuntimeFamily(str, Enum): QWEN_VLLM = "qwen_vllm" QWEN_RUST_CPU = "qwen_rust_cpu" @dataclass class OfflineASRRequest: model_id: str audio_path: str hotwords: str = "" enable_punctuation: bool = True enable_itn: bool = True sample_rate: int = 16000 enable_speaker_diarization: bool = True enable_speaker_identification: bool = True enable_text_cleanup: bool = True word_timestamps: bool = False timestamp_scale: float = 1.0 task_id: Optional[str] = None progress_callback: Optional[Callable[[str, str, int, Optional[dict[str, object]]], None]] = None class RuntimeEngineLease: """Lifecycle wrapper around a pooled engine instance.""" def __init__(self, engine: BaseASREngine, release_callback: Callable[[], None | Awaitable[None]]): self.engine = engine self._release_callback = release_callback self._closed = False async def close(self) -> None: if self._closed: return self._closed = True result = self._release_callback() if asyncio.iscoroutine(result): await result async def __aenter__(self) -> BaseASREngine: return self.engine async def __aexit__(self, exc_type, exc, tb) -> None: await self.close() class RuntimeRouter: """Central backend router for all ASR entrypoints.""" def __init__(self): self._manager = get_model_manager() self._pools: dict[tuple[RuntimeFamily, str], LocalEnginePool[BaseASREngine]] = {} self._shared_engines: dict[tuple[RuntimeFamily, str], BaseASREngine] = {} self._shared_limits: dict[tuple[RuntimeFamily, str], asyncio.Semaphore] = {} self._pool_lock = threading.Lock() self._loaded_model_ids: set[str] = set() def resolve_model_id(self, model_id: Optional[str]) -> str: if model_id: return model_id config = self._manager.get_declared_entry_config() return config.model_id def _resolve_family(self, model_id: str) -> RuntimeFamily: device = detect_device(settings.DEVICE) accelerator = get_accelerator_info() if model_id.startswith("qwen3-asr-"): if accelerator.is_gpu and device.startswith("cuda"): return RuntimeFamily.QWEN_VLLM if device == "cpu" and is_qwenasr_rust_available(): return RuntimeFamily.QWEN_RUST_CPU raise RuntimeError( "Qwen3-ASR is not available on " f"accelerator='{accelerator.vendor}' device='{device}'" ) raise RuntimeError(f"Unsupported runtime model: {model_id}") def _pool_size_for_family(self, family: RuntimeFamily) -> int: if family == RuntimeFamily.QWEN_VLLM: return 1 return settings.QWEN_RUST_CPU_WORKERS def _create_pool(self, family: RuntimeFamily, model_id: str) -> LocalEnginePool[BaseASREngine]: pool_key = (family, model_id) existing = self._pools.get(pool_key) if existing is not None: return existing with self._pool_lock: existing = self._pools.get(pool_key) if existing is not None: return existing pool = LocalEnginePool( size=self._pool_size_for_family(family), factory=lambda: self._manager.create_engine(model_id), ) self._pools[pool_key] = pool self._loaded_model_ids.add(model_id) return pool def _get_shared_engine(self, family: RuntimeFamily, model_id: str) -> tuple[BaseASREngine, asyncio.Semaphore]: runtime_key = (family, model_id) engine = self._shared_engines.get(runtime_key) semaphore = self._shared_limits.get(runtime_key) if engine is not None and semaphore is not None: return engine, semaphore with self._pool_lock: engine = self._shared_engines.get(runtime_key) semaphore = self._shared_limits.get(runtime_key) if engine is None: engine = self._manager.create_engine(model_id) self._shared_engines[runtime_key] = engine self._loaded_model_ids.add(model_id) if semaphore is None: shared_concurrency = max(1, int(settings.QWEN_VLLM_SHARED_CONCURRENCY)) semaphore = asyncio.Semaphore(shared_concurrency) self._shared_limits[runtime_key] = semaphore return engine, semaphore def warmup_model(self, model_id: Optional[str] = None) -> None: resolved_model_id = self.resolve_model_id(model_id) family = self._resolve_family(resolved_model_id) if family == RuntimeFamily.QWEN_VLLM: self._get_shared_engine(family, resolved_model_id) return pool = self._create_pool(family, resolved_model_id) pool.warmup() def get_loaded_model_ids(self) -> list[str]: return sorted(self._loaded_model_ids) def get_memory_usage(self) -> dict[str, object]: memory_info: dict[str, object] = { "model_list": self.get_loaded_model_ids(), "loaded_count": len(self._loaded_model_ids), } accelerator = get_accelerator_info() memory_info["accelerator"] = accelerator.as_dict() if accelerator.is_gpu and torch.cuda.is_available(): memory_info["gpu_memory"] = { "allocated": f"{torch.cuda.memory_allocated() / 1024**3:.2f}GB", "cached": f"{torch.cuda.memory_reserved() / 1024**3:.2f}GB", "max_allocated": f"{torch.cuda.max_memory_allocated() / 1024**3:.2f}GB", } return memory_info async def acquire_engine(self, model_id: Optional[str] = None) -> RuntimeEngineLease: resolved_model_id = self.resolve_model_id(model_id) family = self._resolve_family(resolved_model_id) if family == RuntimeFamily.QWEN_VLLM: engine, semaphore = self._get_shared_engine(family, resolved_model_id) await semaphore.acquire() return RuntimeEngineLease( engine=engine, release_callback=semaphore.release, ) pool = self._create_pool(family, resolved_model_id) engine = await pool.acquire() return RuntimeEngineLease( engine=engine, release_callback=lambda: pool.release(engine), ) async def run_offline(self, request: OfflineASRRequest) -> ASRFullResult: async with await self.acquire_engine(request.model_id) as engine: result = await run_sync( engine.transcribe_long_audio, audio_path=request.audio_path, hotwords=request.hotwords, enable_punctuation=request.enable_punctuation, enable_itn=request.enable_itn, sample_rate=request.sample_rate, enable_speaker_diarization=request.enable_speaker_diarization, enable_speaker_identification=request.enable_speaker_identification, enable_text_cleanup=request.enable_text_cleanup, word_timestamps=request.word_timestamps, timestamp_scale=request.timestamp_scale, task_id=request.task_id, progress_callback=request.progress_callback, ) if ( request.enable_speaker_diarization and request.enable_speaker_identification and settings.SPEAKER_DB_ENABLED ): try: from app.services.speaker_registry import get_speaker_registry_service if request.progress_callback is not None: request.progress_callback( "speaker_matching", "正在匹配已注册声纹库。", 94, None, ) result = await get_speaker_registry_service().apply_registered_speakers( result, threshold=settings.SV_THRESHOLD, ) except Exception: # 数据库/声纹匹配失败不影响原有 ASR 输出,仍保留 CAM++ 的局部说话人编号。 pass return result _runtime_router: Optional[RuntimeRouter] = None _runtime_router_lock = threading.Lock() def get_runtime_router() -> RuntimeRouter: global _runtime_router if _runtime_router is None: with _runtime_router_lock: if _runtime_router is None: _runtime_router = RuntimeRouter() return _runtime_router