import os import re import sys import uuid import subprocess from pathlib import Path import torch from fastapi import FastAPI, HTTPException, Header from fastapi.responses import FileResponse from pydantic import BaseModel from huggingface_hub import snapshot_download MODEL_ID = "beleata74/bg-tts-v5" TMP_DIR = Path("/tmp/books-tts") TMP_DIR.mkdir(parents=True, exist_ok=True) API_KEY = os.getenv("BOOKS_TTS_API_KEY") app = FastAPI(title="Books Bulgarian TTS API") print("Downloading/loading BG-TTS V5 repo...") MODEL_DIR = Path(snapshot_download(repo_id=MODEL_ID)) CHECKPOINT_DIR = MODEL_DIR / "checkpoint" sys.path.insert(0, str(MODEL_DIR)) from tts_v5.model import load_for_inference from tts_v5.tokenizer import TTSTokenizer from tts_v5.codec import CodecV5 from tts_v5.inference import generate from tts_v5.config import CODEC_NUM_CODEBOOKS DEVICE = "cuda" if torch.cuda.is_available() else "cpu" print(f"Loading model on {DEVICE}...") tts_model = load_for_inference(str(CHECKPOINT_DIR), device=DEVICE) tts_tokenizer = TTSTokenizer() tts_codec = CodecV5(device=DEVICE) print("Model loaded.") class TTSRequest(BaseModel): text: str speaker_id: int | None = 1 temperature: float | None = 0.25 top_k: int | None = 50 top_p: float | None = 0.8 format: str | None = "mp3" @app.get("/") def root(): return { "ok": True, "message": "Books Bulgarian TTS API is running", "model": MODEL_ID, "device": DEVICE, } @app.get("/health") def health(): return { "ok": True, "model": MODEL_ID, "device": DEVICE, } def check_auth(authorization: str | None): if not API_KEY: raise HTTPException( status_code=500, detail="BOOKS_TTS_API_KEY is not configured on the Space." ) expected = f"Bearer {API_KEY}" if authorization != expected: raise HTTPException( status_code=401, detail="Invalid or missing Authorization header." ) def split_text(text: str, max_chars: int = 300): text = re.sub(r"\s+", " ", text).strip() if not text: return [] sentences = re.split(r"(?<=[.!?…])\s+", text) chunks = [] current = "" for sentence in sentences: sentence = sentence.strip() if not sentence: continue if len(current) + len(sentence) + 1 <= max_chars: current = f"{current} {sentence}".strip() else: if current: chunks.append(current) if len(sentence) > max_chars: for i in range(0, len(sentence), max_chars): chunks.append(sentence[i:i + max_chars]) current = "" else: current = sentence if current: chunks.append(current) return chunks def generate_wav_for_chunk( text: str, output_path: Path, speaker_id: int, temperature: float, top_k: int, top_p: float, ): tokens = generate( model=tts_model, tokenizer=tts_tokenizer, text=text, speaker_id=speaker_id, max_new_tokens=2000, temperature=temperature, top_k=top_k, top_p=top_p, rep_penalty=1.1, device=DEVICE, ) if tokens is None or len(tokens) == 0: raise RuntimeError("No audio tokens generated.") tokens = tokens[:len(tokens) - len(tokens) % CODEC_NUM_CODEBOOKS] if len(tokens) == 0: raise RuntimeError("Generated token count is invalid.") tts_codec.tokens_to_wav(tokens, str(output_path)) def merge_wavs_to_mp3(wav_paths: list[Path], output_mp3: Path): list_file = TMP_DIR / f"{uuid.uuid4()}-list.txt" with open(list_file, "w", encoding="utf-8") as f: for wav in wav_paths: f.write(f"file '{wav}'\n") cmd = [ "ffmpeg", "-y", "-f", "concat", "-safe", "0", "-i", str(list_file), "-codec:a", "libmp3lame", "-qscale:a", "4", str(output_mp3), ] result = subprocess.run(cmd, capture_output=True, text=True) try: list_file.unlink(missing_ok=True) except Exception: pass if result.returncode != 0: raise RuntimeError(result.stderr or "ffmpeg failed") @app.post("/tts") def tts(payload: TTSRequest, authorization: str | None = Header(default=None)): check_auth(authorization) text = payload.text.strip() if not text: raise HTTPException(status_code=400, detail="Text is empty.") if len(text) > 5000: raise HTTPException( status_code=400, detail="Text is too long. Max 5000 characters for now." ) speaker_id = payload.speaker_id if payload.speaker_id in [0, 1] else 1 request_id = str(uuid.uuid4()) request_dir = TMP_DIR / request_id request_dir.mkdir(parents=True, exist_ok=True) try: max_chars = 300 if speaker_id == 1 else 350 chunks = split_text(text, max_chars=max_chars) if not chunks: raise HTTPException(status_code=400, detail="No valid text chunks found.") wav_paths = [] for index, chunk in enumerate(chunks): wav_path = request_dir / f"chunk-{index}.wav" generate_wav_for_chunk( text=chunk, output_path=wav_path, speaker_id=speaker_id, temperature=payload.temperature or 0.25, top_k=payload.top_k or 50, top_p=payload.top_p or 0.8, ) wav_paths.append(wav_path) output_mp3 = request_dir / "output.mp3" merge_wavs_to_mp3(wav_paths, output_mp3) return FileResponse( path=str(output_mp3), media_type="audio/mpeg", filename="tts.mp3", ) except HTTPException: raise except Exception as e: raise HTTPException( status_code=500, detail=f"TTS generation failed: {str(e)}" )