diff --git a/studio/backend/core/inference/audio_codecs.py b/studio/backend/core/inference/audio_codecs.py index 3a418d921d..bcf3ec2937 100644 --- a/studio/backend/core/inference/audio_codecs.py +++ b/studio/backend/core/inference/audio_codecs.py @@ -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: diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py index 42cce3c4d8..ba45cd7735 100644 --- a/studio/backend/core/inference/llama_cpp.py +++ b/studio/backend/core/inference/llama_cpp.py @@ -779,8 +779,18 @@ 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 + # 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): @@ -1028,3 +1038,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 "")) == 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, + 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", "") + ) diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index 55bd9710a5..13bb35082e 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -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, ) @@ -473,6 +487,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], ) @@ -521,78 +537,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) @@ -714,6 +736,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 868f3aaf33..fa0daea7fd 100644 --- a/studio/backend/routes/models.py +++ b/studio/backend/routes/models.py @@ -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() diff --git a/studio/frontend/src/features/chat/api/chat-adapter.ts b/studio/frontend/src/features/chat/api/chat-adapter.ts index 30a8d97461..6332753b8a 100644 --- a/studio/frontend/src/features/chat/api/chat-adapter.ts +++ b/studio/frontend/src/features/chat/api/chat-adapter.ts @@ -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 {