528 lines
20 KiB
Python
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
|