diff --git a/studio/backend/core/inference/audio_codecs.py b/studio/backend/core/inference/audio_codecs.py index 3a418d921d..f033e6c77c 100644 --- a/studio/backend/core/inference/audio_codecs.py +++ b/studio/backend/core/inference/audio_codecs.py @@ -302,6 +302,22 @@ 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": + return self.decode_snac(torch.tensor([token_ids], dtype = torch.long), device) + elif audio_type == "bicodec": + return self.decode_bicodec(text, device) + elif audio_type == "dac": + return self.decode_dac(text, device) + raise ValueError(f"Cannot decode audio_type: {audio_type}") + # ── Cleanup ────────────────────────────────────────────────── def unload(self) -> None: diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py index 3ed4589dc0..ba32a81fad 100644 --- a/studio/backend/core/inference/llama_cpp.py +++ b/studio/backend/core/inference/llama_cpp.py @@ -981,3 +981,110 @@ 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 "")) == 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": ("{text}<|eot_id|>", [""], True), + "bicodec": ("<|task_tts|><|start_content|>{text}<|end_content|><|start_global_token|>", ["<|im_end|>", ""], 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, + 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, + "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", "")) diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index c7fa9ed720..ea497101a7 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -155,6 +155,16 @@ 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 is not None and _gguf_audio not in ("audio_vlm",) + 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 +174,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, ) @@ -461,6 +474,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], ) @@ -509,78 +524,68 @@ 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, 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, - ), - ) - - 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", - } - ], - } - ) - + wav_bytes, sample_rate = await asyncio.get_event_loop().run_in_executor(None, gen) 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) @@ -702,6 +707,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: diff --git a/studio/backend/routes/models.py b/studio/backend/routes/models.py index 586b18718e..d914ceaa02 100644 --- a/studio/backend/routes/models.py +++ b/studio/backend/routes/models.py @@ -304,6 +304,18 @@ 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() diff --git a/studio/frontend/src/features/chat/api/chat-adapter.ts b/studio/frontend/src/features/chat/api/chat-adapter.ts index 0935728e89..04cd0a5098 100644 --- a/studio/frontend/src/features/chat/api/chat-adapter.ts +++ b/studio/frontend/src/features/chat/api/chat-adapter.ts @@ -79,10 +79,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 {