test/app/services/asr/qwen3_vllm.py

528 lines
20 KiB
Python

# -*- coding: utf-8 -*-
"""Official vLLM adapter for CUDA Qwen3-ASR."""
from __future__ import annotations
import importlib
import importlib.util
import logging
import os
import re
import threading
from dataclasses import dataclass, field
from typing import Any, Optional
import librosa
import numpy as np
from app.core.hotword_resolver import strip_hotword_prompt_leakage
from app.utils.text_processing import normalize_asr_text
from .engines import ASRRawResult, ASRSegmentResult, WordToken
logger = logging.getLogger(__name__)
_DEFAULT_SAMPLE_RATE = 16000
_LANGUAGE_ALIASES = {
"zh": "Chinese",
"zh-cn": "Chinese",
"zh-hans": "Chinese",
"zh-hant": "Chinese",
"cn": "Chinese",
"en": "English",
"en-us": "English",
"en-gb": "English",
"ja": "Japanese",
"jp": "Japanese",
"ko": "Korean",
"yue": "Cantonese",
"fr": "French",
"de": "German",
"es": "Spanish",
"ru": "Russian",
}
_PROMPT_LEAK_PATTERNS = (
re.compile(r"^\s*Transcribe the speech accurately\.\s*", re.IGNORECASE),
re.compile(r"^\s*Transcribe the speech in [A-Za-z\s-]+\.\s*", re.IGNORECASE),
)
def is_vllm_available() -> bool:
"""Return True when the official vLLM runtime is installed."""
return importlib.util.find_spec("vllm") is not None
def _normalize_language_name(language: Optional[str]) -> Optional[str]:
if not language:
return None
normalized = language.strip()
if not normalized:
return None
alias = _LANGUAGE_ALIASES.get(normalized.lower())
if alias:
return alias
if " " in normalized:
return " ".join(part.capitalize() for part in normalized.split())
return normalized.capitalize()
def _load_audio(audio_path: str) -> np.ndarray:
audio, _sample_rate = librosa.load(audio_path, sr=_DEFAULT_SAMPLE_RATE, mono=True)
return audio.astype(np.float32)
def _build_chat_prompt(context: str = "", language: Optional[str] = None) -> str:
instructions: list[str] = []
if language:
instructions.append(f"Transcribe the speech in {language}.")
else:
instructions.append("Transcribe the speech accurately.")
if context.strip():
instructions.append(context.strip())
system_text = " ".join(instructions).strip()
return (
f"<|im_start|>system\n{system_text}<|im_end|>\n"
"<|im_start|>user\n<|audio_start|><|audio_pad|><|audio_end|><|im_end|>\n"
"<|im_start|>assistant\n"
)
def _build_alignment_prompt(tokens: list[str]) -> str:
body = "<timestamp><timestamp>".join(tokens) + "<timestamp><timestamp>"
return f"<|audio_start|><|audio_pad|><|audio_end|>{body}"
def _strip_prompt_leakage(text: str) -> str:
cleaned = text or ""
changed = True
while changed and cleaned:
changed = False
for pattern in _PROMPT_LEAK_PATTERNS:
updated, count = pattern.subn("", cleaned, count=1)
if count:
cleaned = updated
changed = True
return strip_hotword_prompt_leakage(cleaned)
def _sanitize_detected_language(detected: str, fallback: Optional[str]) -> str:
candidate = (detected or "").strip()
if not candidate:
return fallback or ""
candidate = re.sub(r"^[^\w]+", "", candidate)
match = re.search(r"language\s+([A-Za-z][A-Za-z\s-]*)$", candidate, re.IGNORECASE)
if match:
candidate = match.group(1).strip()
normalized = _normalize_language_name(candidate)
return normalized or (fallback or "")
def _parse_asr_output(raw_text: str, language: Optional[str]) -> tuple[str, str]:
text = (raw_text or "").strip()
if "<asr_text>" in text:
left, right = text.split("<asr_text>", 1)
detected = _sanitize_detected_language(left.strip(), language)
return detected, _strip_prompt_leakage(right.strip())
return (language or ""), _strip_prompt_leakage(text)
def _split_alignment_units(text: str) -> list[str]:
if not text:
return []
# Mixed Chinese/English transcripts should not fall back to whitespace-only
# tokenization, otherwise a long CJK sentence with a single embedded English
# word can collapse into one giant alignment unit.
token_pattern = re.compile(
r"[\u4e00-\u9fff]" # CJK ideographs, align per character
r"|[A-Za-z0-9]+(?:['._+-][A-Za-z0-9]+)*" # Latin / alnum words
r"|[^\w\s]", # punctuation and symbols
re.UNICODE,
)
return token_pattern.findall(text)
def _resolve_forced_aligner_gpu_memory_utilization(primary_utilization: float) -> float:
override = (os.getenv("QWEN_FORCE_ALIGNER_GPU_MEMORY_UTILIZATION") or "").strip()
if override:
try:
value = float(override)
if 0.0 < value <= 1.0:
return value
except ValueError:
logger.warning(
"Invalid QWEN_FORCE_ALIGNER_GPU_MEMORY_UTILIZATION=%s, ignoring override",
override,
)
return primary_utilization
@dataclass
class _GeneratedTranscript:
text: str
language: str
@dataclass
class VLLMRealtimeState:
prompt_raw: str
language: str
chunk_size_sec: float
unfixed_chunk_num: int
unfixed_token_num: int
max_new_tokens: int
chunk_id: int = 0
text: str = ""
raw_decoded: str = ""
audio_buffer: np.ndarray = field(default_factory=lambda: np.array([], dtype=np.float32))
audio_accum: np.ndarray = field(default_factory=lambda: np.array([], dtype=np.float32))
class Qwen3VLLMBackend:
"""Thin adapter over official vLLM APIs for Qwen3-ASR."""
def __init__(
self,
model_path: str,
forced_aligner_path: Optional[str],
gpu_memory_utilization: float,
max_inference_batch_size: int,
max_new_tokens: int,
enforce_eager: bool = True,
max_model_len: Optional[int] = None,
tensor_parallel_size: int = 1,
) -> None:
try:
vllm_module = importlib.import_module("vllm")
transformers_module = importlib.import_module("transformers")
except ImportError as exc:
raise RuntimeError(
"CUDA Qwen3-ASR now requires official vLLM with Qwen3 forced aligner support. "
"Install it with: pip install 'vllm[audio]==0.19.0'"
) from exc
self._llm_cls = getattr(vllm_module, "LLM")
self._sampling_params_cls = getattr(vllm_module, "SamplingParams")
self._tokenizer = getattr(transformers_module, "AutoTokenizer").from_pretrained(
model_path,
trust_remote_code=True,
)
llm_kwargs: dict[str, Any] = {
"model": model_path,
"gpu_memory_utilization": gpu_memory_utilization,
"enforce_eager": enforce_eager,
"trust_remote_code": True,
}
if max_model_len is not None:
llm_kwargs["max_model_len"] = max_model_len
if tensor_parallel_size > 1:
llm_kwargs["tensor_parallel_size"] = tensor_parallel_size
self._llm = self._llm_cls(**llm_kwargs)
self._sampling_params = self._sampling_params_cls(
temperature=0.01,
max_tokens=max_new_tokens,
)
self._max_inference_batch_size = max_inference_batch_size
self._gpu_memory_utilization = gpu_memory_utilization
self._enforce_eager = enforce_eager
self._tensor_parallel_size = tensor_parallel_size
self._forced_aligner_path = forced_aligner_path
self._forced_aligner: Any | None = None
self._timestamp_token_id: int | None = None
self._timestamp_segment_time: float | None = None
# Share one backend instance across multiple offline tasks, but serialize
# direct vLLM engine calls to avoid cross-request state corruption/hangs.
self._engine_lock = threading.RLock()
def _get_forced_aligner_gpu_memory_utilization(self) -> float:
configured = _resolve_forced_aligner_gpu_memory_utilization(self._gpu_memory_utilization)
logger.info(
"Resolved forced aligner gpu_memory_utilization=%s (primary=%s)",
configured,
self._gpu_memory_utilization,
)
return configured
def _get_forced_aligner(self) -> Any:
if not self._forced_aligner_path:
raise RuntimeError("word_timestamps requires a configured forced aligner model")
if self._forced_aligner is None:
forced_aligner_gpu_memory_utilization = self._get_forced_aligner_gpu_memory_utilization()
logger.info(
"Loading Qwen3 forced aligner via official vLLM: %s (gpu_memory_utilization=%s)",
self._forced_aligner_path,
forced_aligner_gpu_memory_utilization,
)
self._forced_aligner = self._llm_cls(
model=self._forced_aligner_path,
runner="pooling",
enforce_eager=self._enforce_eager,
gpu_memory_utilization=forced_aligner_gpu_memory_utilization,
tensor_parallel_size=self._tensor_parallel_size,
trust_remote_code=True,
hf_overrides={
"architectures": ["Qwen3ASRForcedAlignerForTokenClassification"],
},
)
llm_engine = getattr(self._forced_aligner, "llm_engine", None)
if llm_engine is None:
raise RuntimeError("Forced aligner did not expose a vLLM engine instance")
config = llm_engine.vllm_config.model_config.hf_config
self._timestamp_token_id = int(config.timestamp_token_id)
self._timestamp_segment_time = float(config.timestamp_segment_time)
return self._forced_aligner
def ensure_forced_aligner_loaded(self) -> None:
if self._forced_aligner_path:
self._get_forced_aligner()
def _run_generate(
self,
audio_items: list[tuple[np.ndarray, str, Optional[str]]],
) -> list[_GeneratedTranscript]:
prompts: list[dict[str, Any]] = []
for audio, context, language in audio_items:
prompts.append(
{
"prompt": _build_chat_prompt(context=context, language=_normalize_language_name(language)),
"multi_modal_data": {"audio": [audio]},
}
)
with self._engine_lock:
outputs = self._llm.generate(
prompts,
sampling_params=self._sampling_params,
use_tqdm=False,
)
transcripts: list[_GeneratedTranscript] = []
for output, (_audio, _context, language) in zip(outputs, audio_items):
raw_text = str(output.outputs[0].text if output.outputs else "")
parsed_language, parsed_text = _parse_asr_output(raw_text, _normalize_language_name(language))
transcripts.append(_GeneratedTranscript(text=parsed_text, language=parsed_language))
return transcripts
def transcribe_text(
self,
audio_path: str,
context: str = "",
language: Optional[str] = None,
enable_itn: bool = False,
) -> str:
transcript = self._run_generate([(_load_audio(audio_path), context, language)])[0]
return normalize_asr_text(transcript.text, enable_itn=enable_itn)
def transcribe_raw(
self,
audio_path: str,
context: str = "",
language: Optional[str] = None,
word_timestamps: bool = False,
enable_itn: bool = False,
) -> ASRRawResult:
audio = _load_audio(audio_path)
transcript = self._run_generate([(audio, context, language)])[0]
text = normalize_asr_text(transcript.text, enable_itn=enable_itn)
if not word_timestamps:
return ASRRawResult(
text=text,
segments=[ASRSegmentResult(text=text, start_time=0.0, end_time=0.0)] if text else [],
)
aligned = self.align_transcript(audio_path=audio_path, text=text, language=language, audio=audio)
word_tokens = [
WordToken(
text=str(item["text"]),
start_time=round(float(item["start_ms"]) / 1000.0, 3),
end_time=round(float(item["end_ms"]) / 1000.0, 3),
)
for item in aligned
]
if not word_tokens:
return ASRRawResult(
text=text,
segments=[ASRSegmentResult(text=text, start_time=0.0, end_time=0.0)] if text else [],
)
return ASRRawResult(
text=text,
segments=[
ASRSegmentResult(
text=text,
start_time=word_tokens[0].start_time,
end_time=word_tokens[-1].end_time,
word_tokens=word_tokens,
)
],
)
def transcribe_batch(
self,
audio_paths: list[str],
context: str = "",
language: Optional[str] = None,
word_timestamps: bool = False,
enable_itn: bool = False,
) -> list[ASRSegmentResult]:
audios = [_load_audio(path) for path in audio_paths]
results: list[ASRSegmentResult] = []
for start in range(0, len(audios), self._max_inference_batch_size):
chunk = audios[start:start + self._max_inference_batch_size]
transcripts = self._run_generate([(audio, context, language) for audio in chunk])
for audio_path, audio, transcript in zip(audio_paths[start:start + len(chunk)], chunk, transcripts):
text = normalize_asr_text(transcript.text, enable_itn=enable_itn)
if not word_timestamps:
results.append(ASRSegmentResult(text=text, start_time=0.0, end_time=0.0))
continue
aligned = self.align_transcript(
audio_path=audio_path,
text=text,
language=language,
audio=audio,
)
word_tokens = [
WordToken(
text=str(item["text"]),
start_time=round(float(item["start_ms"]) / 1000.0, 3),
end_time=round(float(item["end_ms"]) / 1000.0, 3),
)
for item in aligned
]
results.append(
ASRSegmentResult(
text=text,
start_time=word_tokens[0].start_time if word_tokens else 0.0,
end_time=word_tokens[-1].end_time if word_tokens else 0.0,
word_tokens=word_tokens or None,
)
)
return results
def align_transcript(
self,
audio_path: str,
text: str,
language: Optional[str] = None,
audio: Optional[np.ndarray] = None,
) -> list[dict[str, float | str]]:
tokens = _split_alignment_units(text)
if not tokens:
return []
aligner = self._get_forced_aligner()
prompt = _build_alignment_prompt(tokens)
audio_array = audio if audio is not None else _load_audio(audio_path)
with self._engine_lock:
outputs = aligner.encode(
[{"prompt": prompt, "multi_modal_data": {"audio": audio_array}}],
pooling_task="token_classify",
)
output = outputs[0]
logits = output.outputs.data
predictions = logits.argmax(dim=-1) if hasattr(logits, "argmax") else np.argmax(logits, axis=-1)
ts_predictions = [
float(pred.item() if hasattr(pred, "item") else pred) * float(self._timestamp_segment_time or 0.0)
for tid, pred in zip(output.prompt_token_ids, predictions)
if int(tid) == int(self._timestamp_token_id or -1)
]
expected_timestamps = len(tokens) * 2
if len(ts_predictions) < expected_timestamps:
raise RuntimeError(
"Forced aligner returned fewer timestamp predictions than expected: "
f"expected={expected_timestamps}, got={len(ts_predictions)}, tokens={len(tokens)}"
)
aligned: list[dict[str, float | str]] = []
for index, token in enumerate(tokens):
start_ms = ts_predictions[index * 2]
end_ms = ts_predictions[index * 2 + 1]
if end_ms < start_ms:
logger.warning(
"Forced aligner produced reversed timestamps for token=%r: start_ms=%s end_ms=%s",
token,
start_ms,
end_ms,
)
start_ms, end_ms = end_ms, start_ms
aligned.append({"text": token, "start_ms": start_ms, "end_ms": end_ms})
return aligned
def init_streaming_state(
self,
*,
context: str = "",
language: Optional[str] = None,
chunk_size_sec: float = 1.2,
unfixed_chunk_num: int = 2,
unfixed_token_num: int = 5,
max_new_tokens: int = 32,
) -> VLLMRealtimeState:
normalized_language = _normalize_language_name(language) or ""
return VLLMRealtimeState(
prompt_raw=_build_chat_prompt(context=context, language=normalized_language or None),
language=normalized_language,
chunk_size_sec=chunk_size_sec,
unfixed_chunk_num=unfixed_chunk_num,
unfixed_token_num=unfixed_token_num,
max_new_tokens=max_new_tokens,
audio_buffer=np.array([], dtype=np.float32),
audio_accum=np.array([], dtype=np.float32),
)
def _decode_stream(self, state: VLLMRealtimeState) -> VLLMRealtimeState:
prefix = ""
if state.chunk_id >= state.unfixed_chunk_num and state.raw_decoded:
token_ids = self._tokenizer.encode(state.raw_decoded, add_special_tokens=False)
rollback = token_ids[-state.unfixed_token_num:] if state.unfixed_token_num > 0 else []
if rollback:
prefix = self._tokenizer.decode(rollback, skip_special_tokens=False).replace("\ufffd", "")
with self._engine_lock:
output = self._llm.generate(
[
{
"prompt": state.prompt_raw + prefix,
"multi_modal_data": {"audio": [state.audio_accum]},
}
],
sampling_params=self._sampling_params_cls(
temperature=0.01,
max_tokens=state.max_new_tokens,
),
use_tqdm=False,
)[0]
generated = str(output.outputs[0].text if output.outputs else "")
parsed_language, parsed_text = _parse_asr_output(prefix + generated, state.language or None)
state.raw_decoded = prefix + generated
state.text = parsed_text
state.language = parsed_language or state.language
state.chunk_id += 1
return state
def feed_stream(self, pcm: np.ndarray, state: VLLMRealtimeState) -> VLLMRealtimeState:
state.audio_buffer = np.concatenate([state.audio_buffer, pcm.astype(np.float32)])
segment_size = int(max(state.chunk_size_sec, 0.1) * _DEFAULT_SAMPLE_RATE)
while len(state.audio_buffer) >= segment_size:
segment = state.audio_buffer[:segment_size].copy()
state.audio_buffer = state.audio_buffer[segment_size:]
state.audio_accum = np.concatenate([state.audio_accum, segment])
state = self._decode_stream(state)
return state
def finish_stream(self, state: VLLMRealtimeState) -> VLLMRealtimeState:
if len(state.audio_buffer) > 0:
state.audio_accum = np.concatenate([state.audio_accum, state.audio_buffer])
state.audio_buffer = np.array([], dtype=np.float32)
state = self._decode_stream(state)
elif state.chunk_id == 0 and len(state.audio_accum) > 0:
state = self._decode_stream(state)
return state