studio: GGUF TTS audio support (from PR #4318)
Add GGUF TTS audio generation via llama-server. When a GGUF model loads, the backend probes its vocabulary to detect audio codecs (SNAC/BiCodec/DAC/CSM/Whisper). If detected, the codec is pre-loaded and the model is reported as audio to the frontend. During chat, TTS models route to the audio generation path which sends a per-codec prompt to llama-server's /completion endpoint, extracts generated tokens/text, and decodes to WAV using AudioCodecManager. Also strips base64 audio data from prior assistant messages to prevent context overflow. Co-authored-by: Manan Shah <mananshah511@gmail.com>
This commit is contained in:
parent
58523dc4a9
commit
cbb4929139
6 changed files with 17609 additions and 55 deletions
|
|
@ -302,6 +302,28 @@ class AudioCodecManager:
|
|||
waveform = audio.squeeze().cpu().numpy()
|
||||
return _numpy_to_wav_bytes(waveform, 24000), 24000
|
||||
|
||||
def decode(
|
||||
self,
|
||||
audio_type: str,
|
||||
device: str,
|
||||
token_ids: Optional[list] = None,
|
||||
text: Optional[str] = None,
|
||||
) -> Tuple[bytes, int]:
|
||||
"""Unified decode — dispatches to the right codec decoder."""
|
||||
if audio_type == "snac":
|
||||
if not token_ids:
|
||||
raise ValueError("SNAC decoding requires token_ids")
|
||||
return self.decode_snac(torch.tensor([token_ids], dtype = torch.long), device)
|
||||
elif audio_type == "bicodec":
|
||||
if not text:
|
||||
raise ValueError("BiCodec decoding requires text")
|
||||
return self.decode_bicodec(text, device)
|
||||
elif audio_type == "dac":
|
||||
if not text:
|
||||
raise ValueError("DAC decoding requires text")
|
||||
return self.decode_dac(text, device)
|
||||
raise ValueError(f"Cannot decode audio_type: {audio_type}")
|
||||
|
||||
# ── Cleanup ──────────────────────────────────────────────────
|
||||
|
||||
def unload(self) -> None:
|
||||
|
|
|
|||
|
|
@ -872,10 +872,20 @@ class LlamaCppBackend:
|
|||
self._hf_repo = None
|
||||
self._hf_variant = None
|
||||
self._is_vision = False
|
||||
self._is_audio = False
|
||||
self._audio_type = None
|
||||
self._port = None
|
||||
self._healthy = False
|
||||
self._context_length = None
|
||||
self._chat_template = None
|
||||
# Free audio codec GPU memory
|
||||
if LlamaCppBackend._codec_mgr is not None:
|
||||
LlamaCppBackend._codec_mgr.unload()
|
||||
LlamaCppBackend._codec_mgr = None
|
||||
import torch
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
return True
|
||||
|
||||
def _kill_process(self):
|
||||
|
|
@ -1123,3 +1133,149 @@ class LlamaCppBackend:
|
|||
if cancel_event is not None and cancel_event.is_set():
|
||||
return
|
||||
raise
|
||||
|
||||
# ── TTS support ────────────────────────────────────────────
|
||||
|
||||
def detect_audio_type(self) -> Optional[str]:
|
||||
"""Detect audio/TTS codec by probing the loaded model's vocabulary."""
|
||||
if not self.is_loaded:
|
||||
return None
|
||||
try:
|
||||
with httpx.Client(timeout = 10) as client:
|
||||
|
||||
def _detok(tid: int) -> str:
|
||||
r = client.post(
|
||||
f"{self.base_url}/detokenize", json = {"tokens": [tid]}
|
||||
)
|
||||
return r.json().get("content", "") if r.status_code == 200 else ""
|
||||
|
||||
def _tok(text: str) -> list[int]:
|
||||
r = client.post(
|
||||
f"{self.base_url}/tokenize",
|
||||
json = {"content": text, "add_special": False},
|
||||
)
|
||||
return r.json().get("tokens", []) if r.status_code == 200 else []
|
||||
|
||||
# Check codec-specific tokens (not generic ones that may exist in non-audio models)
|
||||
if "<custom_token_" in _detok(128258) and "<custom_token_" in _detok(
|
||||
128259
|
||||
):
|
||||
return "snac"
|
||||
if len(_tok("<|AUDIO|>")) == 1 and len(_tok("<|audio_eos|>")) == 1:
|
||||
return "csm"
|
||||
if len(_tok("<|startoftranscript|>")) == 1:
|
||||
return "whisper"
|
||||
if (
|
||||
len(_tok("<|bicodec_semantic_0|>")) == 1
|
||||
and len(_tok("<|bicodec_global_0|>")) == 1
|
||||
):
|
||||
return "bicodec"
|
||||
if len(_tok("<|c1_0|>")) == 1 and len(_tok("<|c2_0|>")) == 1:
|
||||
return "dac"
|
||||
except Exception as e:
|
||||
logger.debug(f"Audio type detection failed: {e}")
|
||||
return None
|
||||
|
||||
# Prompt format per codec: (template, stop_tokens, needs_token_ids)
|
||||
# Matches prompts in InferenceBackend._generate_snac/bicodec/dac
|
||||
_TTS_PROMPTS = {
|
||||
"snac": (
|
||||
"<custom_token_3>{text}<|eot_id|><custom_token_4>",
|
||||
["<custom_token_2>"],
|
||||
True,
|
||||
),
|
||||
"bicodec": (
|
||||
"<|task_tts|><|start_content|>{text}<|end_content|><|start_global_token|>",
|
||||
["<|im_end|>", "</s>"],
|
||||
False,
|
||||
),
|
||||
"dac": (
|
||||
"<|im_start|>\n<|text_start|>{text}<|text_end|>\n<|audio_start|><|global_features_start|>\n",
|
||||
["<|im_end|>", "<|audio_end|>"],
|
||||
False,
|
||||
),
|
||||
}
|
||||
|
||||
_codec_mgr = None # Shared AudioCodecManager instance
|
||||
|
||||
def init_audio_codec(self, audio_type: str) -> None:
|
||||
"""Load the audio codec at model load time (mirrors non-GGUF path)."""
|
||||
import torch
|
||||
from core.inference.audio_codecs import AudioCodecManager
|
||||
|
||||
if LlamaCppBackend._codec_mgr is None:
|
||||
LlamaCppBackend._codec_mgr = AudioCodecManager()
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
model_repo_path = None
|
||||
|
||||
# BiCodec needs a repo with BiCodec/ weights — download canonical SparkTTS
|
||||
if audio_type == "bicodec":
|
||||
from huggingface_hub import snapshot_download
|
||||
import os
|
||||
|
||||
repo_path = snapshot_download(
|
||||
"unsloth/Spark-TTS-0.5B", local_dir = "Spark-TTS-0.5B"
|
||||
)
|
||||
model_repo_path = os.path.abspath(repo_path)
|
||||
|
||||
LlamaCppBackend._codec_mgr.load_codec(
|
||||
audio_type, device, model_repo_path = model_repo_path
|
||||
)
|
||||
logger.info(f"Loaded audio codec for GGUF TTS: {audio_type}")
|
||||
|
||||
def generate_audio_response(
|
||||
self,
|
||||
text: str,
|
||||
audio_type: 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,
|
||||
) -> tuple:
|
||||
"""
|
||||
Generate TTS audio via llama-server /completion + codec decoding.
|
||||
Returns (wav_bytes, sample_rate).
|
||||
"""
|
||||
if audio_type not in self._TTS_PROMPTS:
|
||||
raise RuntimeError(f"GGUF TTS does not support '{audio_type}' codec.")
|
||||
|
||||
tpl, stop, need_ids = self._TTS_PROMPTS[audio_type]
|
||||
|
||||
payload: dict = {
|
||||
"prompt": tpl.format(text = text),
|
||||
"stream": False,
|
||||
"n_predict": max_new_tokens,
|
||||
"temperature": temperature,
|
||||
"top_p": top_p,
|
||||
"top_k": top_k if top_k >= 0 else 0,
|
||||
"min_p": min_p,
|
||||
"repeat_penalty": repetition_penalty,
|
||||
}
|
||||
if stop:
|
||||
payload["stop"] = stop
|
||||
if need_ids:
|
||||
payload["n_probs"] = 1
|
||||
|
||||
with httpx.Client(timeout = httpx.Timeout(300, connect = 10)) as client:
|
||||
resp = client.post(f"{self.base_url}/completion", json = payload)
|
||||
if resp.status_code != 200:
|
||||
raise RuntimeError(
|
||||
f"llama-server returned {resp.status_code}: {resp.text}"
|
||||
)
|
||||
|
||||
data = resp.json()
|
||||
token_ids = (
|
||||
[p["id"] for p in data.get("completion_probabilities", []) if "id" in p]
|
||||
if need_ids
|
||||
else None
|
||||
)
|
||||
|
||||
import torch
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
return LlamaCppBackend._codec_mgr.decode(
|
||||
audio_type, device, token_ids = token_ids, text = data.get("content", "")
|
||||
)
|
||||
|
|
|
|||
|
|
@ -155,6 +155,17 @@ async def load_model(
|
|||
|
||||
logger.info(f"Loaded GGUF model via llama-server: {config.identifier}")
|
||||
|
||||
# Detect TTS audio by probing the loaded model's vocabulary
|
||||
from utils.models import is_audio_input_type
|
||||
|
||||
_gguf_audio = llama_backend.detect_audio_type()
|
||||
_gguf_is_audio = _gguf_audio in ("snac", "bicodec", "dac")
|
||||
llama_backend._is_audio = _gguf_is_audio
|
||||
llama_backend._audio_type = _gguf_audio
|
||||
if _gguf_is_audio:
|
||||
logger.info(f"GGUF model detected as audio: audio_type={_gguf_audio}")
|
||||
await asyncio.to_thread(llama_backend.init_audio_codec, _gguf_audio)
|
||||
|
||||
inference_config = load_inference_config(config.identifier)
|
||||
|
||||
return LoadResponse(
|
||||
|
|
@ -164,6 +175,9 @@ async def load_model(
|
|||
is_vision = config.is_vision,
|
||||
is_lora = False,
|
||||
is_gguf = True,
|
||||
is_audio = _gguf_is_audio,
|
||||
audio_type = _gguf_audio,
|
||||
has_audio_input = is_audio_input_type(_gguf_audio),
|
||||
inference = inference_config,
|
||||
context_length = llama_backend.context_length,
|
||||
)
|
||||
|
|
@ -474,6 +488,8 @@ async def get_status(
|
|||
is_vision = llama_backend.is_vision,
|
||||
is_gguf = True,
|
||||
gguf_variant = llama_backend.hf_variant,
|
||||
is_audio = getattr(llama_backend, "_is_audio", False),
|
||||
audio_type = getattr(llama_backend, "_audio_type", None),
|
||||
loading = [],
|
||||
loaded = [llama_backend.model_identifier],
|
||||
)
|
||||
|
|
@ -522,78 +538,84 @@ async def generate_audio(
|
|||
"""
|
||||
Generate audio (TTS) from the latest user message.
|
||||
Returns a JSON response with base64-encoded WAV audio.
|
||||
Only works when an audio model is loaded.
|
||||
Works with both GGUF (llama-server) and Unsloth/transformers backends.
|
||||
"""
|
||||
import base64
|
||||
|
||||
backend = get_inference_backend()
|
||||
if not backend.active_model_name:
|
||||
raise HTTPException(status_code = 400, detail = "No model loaded.")
|
||||
|
||||
model_info = backend.models.get(backend.active_model_name, {})
|
||||
if not model_info.get("is_audio"):
|
||||
raise HTTPException(
|
||||
status_code = 400, detail = "Active model is not an audio model."
|
||||
)
|
||||
|
||||
# Extract text from the last user message
|
||||
_, chat_messages, _ = _extract_content_parts(payload.messages)
|
||||
if not chat_messages:
|
||||
raise HTTPException(status_code = 400, detail = "No messages provided.")
|
||||
|
||||
last_user_msg = next(
|
||||
(m for m in reversed(chat_messages) if m["role"] == "user"), None
|
||||
)
|
||||
if not last_user_msg:
|
||||
raise HTTPException(status_code = 400, detail = "No user message found.")
|
||||
|
||||
text = last_user_msg["content"]
|
||||
|
||||
# Pick backend — both return (wav_bytes, sample_rate)
|
||||
llama_backend = get_llama_cpp_backend()
|
||||
if llama_backend.is_loaded and getattr(llama_backend, "_is_audio", False):
|
||||
model_name = llama_backend.model_identifier
|
||||
gen = lambda: llama_backend.generate_audio_response(
|
||||
text = text,
|
||||
audio_type = llama_backend._audio_type,
|
||||
temperature = payload.temperature,
|
||||
top_p = payload.top_p,
|
||||
top_k = payload.top_k,
|
||||
min_p = payload.min_p,
|
||||
max_new_tokens = payload.max_tokens or 2048,
|
||||
repetition_penalty = payload.repetition_penalty,
|
||||
)
|
||||
else:
|
||||
backend = get_inference_backend()
|
||||
if not backend.active_model_name:
|
||||
raise HTTPException(status_code = 400, detail = "No model loaded.")
|
||||
model_info = backend.models.get(backend.active_model_name, {})
|
||||
if not model_info.get("is_audio"):
|
||||
raise HTTPException(
|
||||
status_code = 400, detail = "Active model is not an audio model."
|
||||
)
|
||||
model_name = backend.active_model_name
|
||||
gen = lambda: backend.generate_audio_response(
|
||||
text = text,
|
||||
temperature = payload.temperature,
|
||||
top_p = payload.top_p,
|
||||
top_k = payload.top_k,
|
||||
min_p = payload.min_p,
|
||||
max_new_tokens = payload.max_tokens or 2048,
|
||||
repetition_penalty = payload.repetition_penalty,
|
||||
use_adapter = payload.use_adapter,
|
||||
)
|
||||
|
||||
try:
|
||||
wav_bytes, sample_rate = await asyncio.get_event_loop().run_in_executor(
|
||||
None,
|
||||
lambda: backend.generate_audio_response(
|
||||
text = text,
|
||||
temperature = payload.temperature,
|
||||
top_p = payload.top_p,
|
||||
top_k = payload.top_k,
|
||||
min_p = payload.min_p,
|
||||
max_new_tokens = payload.max_tokens or 2048,
|
||||
repetition_penalty = payload.repetition_penalty,
|
||||
use_adapter = payload.use_adapter,
|
||||
),
|
||||
None, gen
|
||||
)
|
||||
|
||||
audio_b64 = base64.b64encode(wav_bytes).decode("ascii")
|
||||
completion_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||
|
||||
return JSONResponse(
|
||||
content = {
|
||||
"id": completion_id,
|
||||
"object": "chat.completion.audio",
|
||||
"model": backend.active_model_name,
|
||||
"audio": {
|
||||
"data": audio_b64,
|
||||
"format": "wav",
|
||||
"sample_rate": sample_rate,
|
||||
},
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": f'[Generated audio from: "{text[:100]}"]',
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Audio generation error: {e}", exc_info = True)
|
||||
raise HTTPException(status_code = 500, detail = str(e))
|
||||
|
||||
audio_b64 = base64.b64encode(wav_bytes).decode("ascii")
|
||||
return JSONResponse(
|
||||
content = {
|
||||
"id": f"chatcmpl-{uuid.uuid4().hex[:12]}",
|
||||
"object": "chat.completion.audio",
|
||||
"model": model_name,
|
||||
"audio": {"data": audio_b64, "format": "wav", "sample_rate": sample_rate},
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": f'[Generated audio from: "{text[:100]}"]',
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
# =====================================================================
|
||||
# OpenAI-Compatible Chat Completions (/chat/completions)
|
||||
|
|
@ -715,6 +737,8 @@ async def openai_chat_completions(
|
|||
# ── Determine which backend is active ─────────────────────
|
||||
if using_gguf:
|
||||
model_name = llama_backend.model_identifier or payload.model
|
||||
if getattr(llama_backend, "_is_audio", False):
|
||||
return await generate_audio(payload, request)
|
||||
else:
|
||||
backend = get_inference_backend()
|
||||
if not backend.active_model_name:
|
||||
|
|
|
|||
|
|
@ -304,6 +304,22 @@ async def list_models(
|
|||
)
|
||||
loaded_models.append(model_info)
|
||||
|
||||
# Include active GGUF model (loaded via llama-server)
|
||||
from routes.inference import get_llama_cpp_backend
|
||||
|
||||
llama_backend = get_llama_cpp_backend()
|
||||
if llama_backend.is_loaded and llama_backend.model_identifier:
|
||||
loaded_models.append(
|
||||
ModelDetails(
|
||||
id = llama_backend.model_identifier,
|
||||
name = llama_backend.model_identifier.split("/")[-1],
|
||||
is_gguf = True,
|
||||
is_vision = llama_backend.is_vision,
|
||||
is_audio = getattr(llama_backend, "_is_audio", False),
|
||||
audio_type = getattr(llama_backend, "_audio_type", None),
|
||||
)
|
||||
)
|
||||
|
||||
# Combine default and loaded models
|
||||
all_models = []
|
||||
seen_ids = set()
|
||||
|
|
|
|||
17329
studio/frontend/package-lock.json
generated
Normal file
17329
studio/frontend/package-lock.json
generated
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -86,10 +86,17 @@ function toOpenAIMessage(message: RunMessage): {
|
|||
return null;
|
||||
}
|
||||
|
||||
return {
|
||||
role: message.role,
|
||||
content: collectTextParts(message).join("\n"),
|
||||
};
|
||||
let content = collectTextParts(message).join("\n");
|
||||
// Strip inline audio base64 from prior assistant messages to avoid
|
||||
// inflating token counts (e.g. audio-player responses with embedded WAV).
|
||||
if (message.role === "assistant") {
|
||||
content = content.replace(
|
||||
/data:audio\/[a-z0-9.+-]+;base64,[A-Za-z0-9+/=]+/g,
|
||||
"[audio]",
|
||||
);
|
||||
}
|
||||
|
||||
return { role: message.role, content };
|
||||
}
|
||||
|
||||
function extractImageBase64(input: string): string | undefined {
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue