feat: route audio inference (TTS, ASR, Whisper) through orchestrator/worker subprocess
This commit is contained in:
parent
f4393ed3e5
commit
7ee81dd7df
2 changed files with 315 additions and 0 deletions
|
|
@ -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)
|
||||
# ------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue