import io import os import logging import numpy as np from typing import Optional import torch from ..config import settings logger = logging.getLogger(__name__) _model = None _tokenizer = None _use_http = True def _load_http_mode(): global _use_http import requests try: resp = requests.get( settings.tts_api_url.replace("/v1/audio/speech", "/health"), timeout=5 ) if resp.ok: _use_http = True logger.info("TTS sidecar available at %s", settings.tts_api_url) return True except Exception: pass _use_http = False return False def _load_model(): global _model, _tokenizer if _model is not None: return logger.info("Loading Higgs TTS 3 on %s", settings.tts_device) from transformers import AutoModel, AutoTokenizer _model = AutoModel.from_pretrained( settings.tts_model_id, torch_dtype=torch.bfloat16, trust_remote_code=True, device_map=settings.tts_device, ) _model.eval() _tokenizer = AutoTokenizer.from_pretrained(settings.tts_model_id) logger.info("Higgs TTS 3 loaded successfully") def synthesize_segment( text: str, reference_audio_path: Optional[str] = None, reference_text: Optional[str] = None, speaker_label: Optional[str] = None, ) -> bytes: if not text.strip(): return b"" input_text = text if speaker_label: input_text = f"<|emotion:contentment|><|prosody:expressive_low|>{text}" if _use_http: try: return _synthesize_http(input_text, reference_audio_path, reference_text) except Exception as e: logger.warning("HTTP TTS failed, falling back to direct: %s", e) return _synthesize_direct(input_text, reference_audio_path, reference_text) def _synthesize_http( text: str, reference_audio_path: Optional[str] = None, reference_text: Optional[str] = None, ) -> bytes: import requests payload = { "input": text, "model": settings.tts_model_id, "temperature": 0.8, "top_k": 50, "max_new_tokens": 2048, } if reference_audio_path and reference_text: import base64 with open(reference_audio_path, "rb") as f: audio_b64 = base64.b64encode(f.read()).decode() payload["references"] = [ { "audio": audio_b64, "text": reference_text, } ] resp = requests.post(settings.tts_api_url, json=payload, timeout=120) resp.raise_for_status() return resp.content def _synthesize_direct( text: str, reference_audio_path: Optional[str] = None, reference_text: Optional[str] = None, ) -> bytes: _load_model() from boson_multimodal.data_types import ChatMLSample, Message from boson_multimodal.dataset.chatml_dataset import ( ChatMLDatasetSample, prepare_chatml_sample_qwen, ) from boson_multimodal.data_collator.higgs_audio_collator import ( HiggsAudioSampleCollator, ) from functools import partial from dataclasses import asdict prompt = f"Generate speech for the following text. Output only the spoken words." messages = [Message(role="user", content=[prompt])] chatml = ChatMLSample(messages=messages) prep_fn = partial(prepare_chatml_sample_qwen, enable_thinking=False) input_tokens, _, _, _ = prep_fn(chatml, _tokenizer, add_generation_prompt=True) ref_audio = None if reference_audio_path: import soundfile as sf ref_audio, ref_sr = sf.read(reference_audio_path) if ref_sr != 16000: import librosa ref_audio = librosa.resample(ref_audio, orig_sr=ref_sr, target_sr=16000) ref_audio = torch.tensor(ref_audio, dtype=torch.float32) sample = ChatMLDatasetSample( input_ids=torch.LongTensor(input_tokens), label_ids=None, audio_ids_concat=None, audio_ids_start=None, audio_waveforms_concat=ref_audio if ref_audio is not None else torch.zeros(1), audio_waveforms_start=torch.tensor([0]), audio_sample_rate=torch.tensor([16000]), audio_speaker_indices=torch.tensor([0]), ) collator = HiggsAudioSampleCollator( whisper_processor=None, audio_in_token_id=getattr(_model.config, "audio_in_token_idx", 0), audio_out_token_id=getattr(_model.config, "audio_out_token_idx", 1), audio_stream_bos_id=getattr(_model.config, "audio_stream_bos_id", 2), audio_stream_eos_id=getattr(_model.config, "audio_stream_eos_id", 3), encode_whisper_embed=getattr(_model.config, "encode_whisper_embed", False), pad_token_id=_tokenizer.pad_token_id or 0, return_audio_in_tokens=getattr(_model.config, "encode_audio_in_tokens", False), use_delay_pattern=getattr(_model.config, "use_delay_pattern", True), round_to=1, audio_num_codebooks=getattr(_model.config, "audio_num_codebooks", 8), chunk_size_seconds=getattr(_model.config, "chunk_size_seconds", 30), encoder_padding_method=getattr( _model.config, "encoder_padding_method", "max_length" ), ) batch = asdict(collator([sample])) device = next(_model.parameters()).device batch = { k: v.to(device).contiguous() if isinstance(v, torch.Tensor) else v for k, v in batch.items() } with torch.inference_mode(): outputs = _model.generate( **batch, max_new_tokens=2048, temperature=0.8, top_k=50, do_sample=True, use_cache=True, ) output_ids = outputs[0] if isinstance(outputs, tuple) else outputs if hasattr(_tokenizer, "decode_audio"): audio_data = _tokenizer.decode_audio(output_ids) if isinstance(audio_data, bytes): return audio_data return audio_data.tobytes() if hasattr(_model, "decode_audio"): audio_data = _model.decode_audio(output_ids) if isinstance(audio_data, bytes): return audio_data if isinstance(audio_data, torch.Tensor): audio_np = audio_data.cpu().float().numpy() return _numpy_to_wav_bytes(audio_np, 24000) if isinstance(output_ids, torch.Tensor): audio_np = output_ids.cpu().float().numpy() if audio_np.ndim > 1: audio_np = audio_np.flatten() return _numpy_to_wav_bytes(audio_np, 24000) return b"" def _numpy_to_wav_bytes(audio: np.ndarray, sample_rate: int = 24000) -> bytes: import wave import struct audio = np.clip(audio, -1.0, 1.0) audio_int16 = (audio * 32767).astype(np.int16) buf = io.BytesIO() with wave.open(buf, "wb") as wf: wf.setnchannels(1) wf.setsampwidth(2) wf.setframerate(sample_rate) wf.writeframes(audio_int16.tobytes()) return buf.getvalue() def _bytes_to_numpy( audio_bytes: bytes, target_sr: int = 24000 ) -> tuple[np.ndarray, int]: import soundfile as sf with io.BytesIO(audio_bytes) as buf: data, sr = sf.read(buf) if len(data.shape) > 1: data = data.mean(axis=1) if sr != target_sr: import librosa data = librosa.resample(data, orig_sr=sr, target_sr=target_sr) return data.astype(np.float32), target_sr def unload_model(): global _model, _tokenizer _model = None _tokenizer = None if torch.cuda.is_available(): torch.cuda.empty_cache()