feat: route audio inference (TTS, ASR, Whisper) through orchestrator/worker subprocess

This commit is contained in:
Roland Tannous 2026-03-08 18:25:27 +00:00
commit 7ee81dd7df
2 changed files with 315 additions and 0 deletions

View file

@ -312,6 +312,9 @@ class InferenceOrchestrator:
"is_vision": model_info.get("is_vision", False),
"is_lora": model_info.get("is_lora", False),
"display_name": model_info.get("display_name", model_name),
"is_audio": model_info.get("is_audio", False),
"audio_type": model_info.get("audio_type"),
"has_audio_input": model_info.get("has_audio_input", False),
}
self.loading_models.discard(model_name)
logger.info("Model '%s' loaded successfully in subprocess", model_name)
@ -545,6 +548,203 @@ class InferenceOrchestrator:
except RuntimeError:
pass
# ------------------------------------------------------------------
# Audio generation — TTS, ASR, audio input
# ------------------------------------------------------------------
def generate_audio_response(
self,
text: str,
temperature: float = 0.6,
top_p: float = 0.95,
top_k: int = 50,
min_p: float = 0.0,
max_new_tokens: int = 2048,
repetition_penalty: float = 1.1,
use_adapter: Optional[Union[bool, str]] = None,
) -> Tuple[bytes, int]:
"""Generate TTS audio. Returns (wav_bytes, sample_rate).
Blocking sends command and waits for the complete audio response.
"""
if not self._ensure_subprocess_alive():
raise RuntimeError("Inference subprocess is not running")
if not self.active_model_name:
raise RuntimeError("No active model")
import uuid
request_id = str(uuid.uuid4())
cmd = {
"type": "generate_audio",
"request_id": request_id,
"text": text,
"temperature": temperature,
"top_p": top_p,
"top_k": top_k,
"min_p": min_p,
"max_new_tokens": max_new_tokens,
"repetition_penalty": repetition_penalty,
}
if use_adapter is not None:
cmd["use_adapter"] = use_adapter
self._send_cmd(cmd)
# Wait for audio_done or audio_error
deadline = time.monotonic() + 120.0
while time.monotonic() < deadline:
remaining = max(0.1, deadline - time.monotonic())
resp = self._read_resp(timeout=min(remaining, 1.0))
if resp is None:
if not self._ensure_subprocess_alive():
raise RuntimeError("Inference subprocess crashed during audio generation")
continue
rtype = resp.get("type", "")
if rtype == "audio_done":
wav_bytes = base64.b64decode(resp["wav_base64"])
sample_rate = resp["sample_rate"]
return wav_bytes, sample_rate
if rtype == "audio_error":
raise RuntimeError(resp.get("error", "Audio generation failed"))
if rtype == "error":
raise RuntimeError(resp.get("error", "Unknown error"))
if rtype == "status":
continue
raise RuntimeError("Timeout waiting for audio generation (120s)")
def generate_whisper_response(
self,
audio_array,
cancel_event=None,
) -> Generator[str, None, None]:
"""Whisper ASR — sends audio to subprocess, yields text."""
yield from self._generate_audio_input_inner(
audio_array=audio_array,
audio_type="whisper",
messages=[],
system_prompt="",
cancel_event=cancel_event,
)
def generate_audio_input_response(
self,
messages,
system_prompt,
audio_array,
temperature: float = 0.7,
top_p: float = 0.9,
top_k: int = 40,
min_p: float = 0.0,
max_new_tokens: int = 512,
repetition_penalty: float = 1.1,
cancel_event=None,
) -> Generator[str, None, None]:
"""Audio input generation (e.g. Gemma 3n) — streams text tokens."""
yield from self._generate_audio_input_inner(
audio_array=audio_array,
audio_type=None, # worker will use generate_audio_input_response
messages=messages,
system_prompt=system_prompt,
temperature=temperature,
top_p=top_p,
top_k=top_k,
min_p=min_p,
max_new_tokens=max_new_tokens,
repetition_penalty=repetition_penalty,
cancel_event=cancel_event,
)
def _generate_audio_input_inner(
self,
audio_array,
audio_type: Optional[str] = None,
messages: list = None,
system_prompt: str = "",
temperature: float = 0.7,
top_p: float = 0.9,
top_k: int = 40,
min_p: float = 0.0,
max_new_tokens: int = 512,
repetition_penalty: float = 1.1,
cancel_event=None,
) -> Generator[str, None, None]:
"""Shared inner logic for audio input generation (Whisper + ASR)."""
if not self._ensure_subprocess_alive():
yield "Error: Inference subprocess is not running"
return
if not self.active_model_name:
yield "Error: No active model"
return
with self._gen_lock:
import uuid
request_id = str(uuid.uuid4())
# Convert numpy array to list for mp.Queue serialization
audio_data = audio_array.tolist() if hasattr(audio_array, 'tolist') else list(audio_array)
cmd = {
"type": "generate_audio_input",
"request_id": request_id,
"audio_data": audio_data,
"audio_type": audio_type,
"messages": messages or [],
"system_prompt": system_prompt,
"temperature": temperature,
"top_p": top_p,
"top_k": top_k,
"min_p": min_p,
"max_new_tokens": max_new_tokens,
"repetition_penalty": repetition_penalty,
}
try:
self._send_cmd(cmd)
except RuntimeError as exc:
yield f"Error: {exc}"
return
# Yield tokens — same pattern as _generate_locked
while True:
resp = self._read_resp(timeout=30.0)
if resp is None:
if not self._ensure_subprocess_alive():
yield "Error: Inference subprocess crashed during audio input generation"
return
continue
rtype = resp.get("type", "")
if rtype == "status":
continue
if rtype == "error" and not resp.get("request_id"):
yield f"Error: {resp.get('error', 'Unknown error')}"
return
if rtype == "token":
if cancel_event is not None and cancel_event.is_set():
self._cancel_generation()
self._drain_until_gen_done(timeout=5.0)
return
yield resp.get("text", "")
elif rtype == "gen_done":
return
elif rtype == "gen_error":
yield f"Error: {resp.get('error', 'Unknown error')}"
return
# ------------------------------------------------------------------
# Local helpers (no subprocess needed)
# ------------------------------------------------------------------

View file

@ -165,6 +165,9 @@ def _handle_load(backend, config: dict, resp_queue: Any) -> None:
"is_vision": mc.is_vision,
"is_lora": mc.is_lora,
"is_gguf": False,
"is_audio": getattr(mc, "is_audio", False),
"audio_type": getattr(mc, "audio_type", None),
"has_audio_input": getattr(mc, "has_audio_input", False),
}
_send_response(resp_queue, {
"type": "loaded",
@ -267,6 +270,110 @@ def _handle_generate(
})
def _handle_generate_audio(
backend,
cmd: dict,
resp_queue: Any,
) -> None:
"""Handle TTS audio generation — returns WAV bytes + sample_rate."""
request_id = cmd.get("request_id", "")
try:
wav_bytes, sample_rate = backend.generate_audio_response(
text=cmd["text"],
temperature=cmd.get("temperature", 0.6),
top_p=cmd.get("top_p", 0.95),
top_k=cmd.get("top_k", 50),
min_p=cmd.get("min_p", 0.0),
max_new_tokens=cmd.get("max_new_tokens", 2048),
repetition_penalty=cmd.get("repetition_penalty", 1.1),
use_adapter=cmd.get("use_adapter"),
)
# Send WAV bytes as base64 (bytes can't go through mp.Queue directly)
_send_response(resp_queue, {
"type": "audio_done",
"request_id": request_id,
"wav_base64": base64.b64encode(wav_bytes).decode("ascii"),
"sample_rate": sample_rate,
"ts": time.time(),
})
except Exception as exc:
logger.error("Audio generation error: %s", exc, exc_info=True)
_send_response(resp_queue, {
"type": "audio_error",
"request_id": request_id,
"error": str(exc),
"stack": traceback.format_exc(limit=20),
"ts": time.time(),
})
def _handle_generate_audio_input(
backend,
cmd: dict,
resp_queue: Any,
cancel_event,
) -> None:
"""Handle audio input generation (ASR/Whisper) — streams text tokens back."""
request_id = cmd.get("request_id", "")
try:
import numpy as np
# Decode audio array from list (numpy arrays can't go through mp.Queue)
audio_array = np.array(cmd["audio_data"], dtype=np.float32)
audio_type = cmd.get("audio_type")
if audio_type == "whisper":
generator = backend.generate_whisper_response(
audio_array=audio_array,
cancel_event=cancel_event,
)
else:
generator = backend.generate_audio_input_response(
messages=cmd.get("messages", []),
system_prompt=cmd.get("system_prompt", ""),
audio_array=audio_array,
temperature=cmd.get("temperature", 0.7),
top_p=cmd.get("top_p", 0.9),
top_k=cmd.get("top_k", 40),
min_p=cmd.get("min_p", 0.0),
max_new_tokens=cmd.get("max_new_tokens", 512),
repetition_penalty=cmd.get("repetition_penalty", 1.1),
cancel_event=cancel_event,
)
for text_chunk in generator:
if cancel_event.is_set():
logger.info("Audio input generation cancelled for request %s", request_id)
break
_send_response(resp_queue, {
"type": "token",
"request_id": request_id,
"text": text_chunk,
"ts": time.time(),
})
_send_response(resp_queue, {
"type": "gen_done",
"request_id": request_id,
"ts": time.time(),
})
except Exception as exc:
logger.error("Audio input generation error: %s", exc, exc_info=True)
_send_response(resp_queue, {
"type": "gen_error",
"request_id": request_id,
"error": str(exc),
"stack": traceback.format_exc(limit=20),
"ts": time.time(),
})
def _handle_unload(backend, cmd: dict, resp_queue: Any) -> None:
"""Handle an unload command."""
model_name = cmd.get("model_name", "")
@ -414,6 +521,14 @@ def run_inference_process(
backend.unload_model(backend.active_model_name)
_handle_load(backend, cmd, resp_queue)
elif cmd_type == "generate_audio":
cancel_event.clear()
_handle_generate_audio(backend, cmd, resp_queue)
elif cmd_type == "generate_audio_input":
cancel_event.clear()
_handle_generate_audio_input(backend, cmd, resp_queue, cancel_event)
elif cmd_type == "unload":
_handle_unload(backend, cmd, resp_queue)