CoolFace
Apppublic

ebitlogix/Parler_TTS_API

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
api.py457 linesDownload Raw Back to root
1from fastapi import FastAPI, HTTPException, WebSocket, WebSocketDisconnect2from fastapi.responses import StreamingResponse3import json4import torch5import numpy as np6import re7from io import BytesIO8import soundfile as sf9from pydantic import BaseModel10import os11from huggingface_hub import login12from parler_tts import ParlerTTSForConditionalGeneration, ParlerTTSStreamer13from transformers import AutoTokenizer14from threading import Thread15import queue16 17# Authenticate with HuggingFace if token is available18hf_token = os.getenv("HF_TOKEN")19if hf_token:20    login(token=hf_token)21 22# Try to import spaces for HF Spaces deployment23try:24    import spaces25    HAS_SPACES = True26except ImportError:27    HAS_SPACES = False28    class _NoOpSpaces:29        def GPU(self, *args, **kwargs):30            def decorator(fn):31                return fn32            return decorator33    spaces = _NoOpSpaces()34 35# --- Model Loading ---36 37MODEL_ID = "ai4bharat/indic-parler-tts"38DEVICE = "cuda" if torch.cuda.is_available() else "cpu"39 40print(f"Using device: {DEVICE}")41if DEVICE == "cuda":42    print(f"GPU: {torch.cuda.get_device_name(0)}")43 44print("Loading Indic Parler-TTS model...")45model = ParlerTTSForConditionalGeneration.from_pretrained(MODEL_ID).to(DEVICE)46 47# Optimize model for inference48if DEVICE == "cuda":49    model = model.half()  # Use half precision (fp16) for faster inference50    model.eval()51 52tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)53description_tokenizer = AutoTokenizer.from_pretrained(model.config.text_encoder._name_or_path)54SAMPLE_RATE = model.config.sampling_rate55 56# Disable gradients for inference57torch.set_grad_enabled(False)58 59print("Model loaded and optimized!")60 61# Named speakers62SPEAKERS = {63    "Divya": "Divya",64    "Rani": "Rani",65    "Rohit": "Rohit",66    "Aman": "Aman",67    "Generic Female": "",68    "Generic Male": "",69}70 71app = FastAPI(title="Parler TTS API", version="1.0")72 73 74class TTSRequest(BaseModel):75    text: str76    speaker: str = "Divya"77    pitch: str = "Moderate"78    rate: str = "Moderate"79    temperature: float = 0.880    do_sample: bool = True81 82 83def build_description(speaker_name, gender, pitch, rate):84    """Build voice description prompt."""85    if speaker_name:86        return (87            f"{speaker_name}'s voice delivers a slightly expressive speech "88            f"with a {pitch.lower()} pitch and a {rate.lower()} speaking rate. "89            f"The recording is of very high quality, with the speaker's voice sounding clear "90            f"and very close up. Very clear audio."91        )92    else:93        return (94            f"A {gender.lower()} speaker delivers a slightly expressive and clear speech "95            f"with a {pitch.lower()} pitch and a {rate.lower()} speaking rate. "96            f"The recording is of very high quality, with the speaker's voice sounding clear "97            f"and very close up. Very clear audio."98        )99 100 101def split_sentences(text):102    """Split Urdu text into sentences."""103    sentences = re.split(r'[۔।\.\!\?]+', text)104    return [s.strip() for s in sentences if s.strip()]105 106 107def clean_urdu_text(text):108    """Minimal text cleaning - preserve content."""109    text = re.sub(r'\s+', ' ', text).strip()110    if text and text[-1] not in '۔.!?،':111        text += '۔'112    return text113 114 115@spaces.GPU()116def generate_speech_internal(text, speaker, pitch, rate, temperature, do_sample):117    """Internal function for speech generation - sentence by sentence with fp16."""118    if not text.strip():119        return None120 121    try:122        text = clean_urdu_text(text)123        speaker_name = SPEAKERS.get(speaker, "")124        gender = "female" if "Female" in speaker or speaker in ["Divya", "Rani"] else "male"125        description = build_description(speaker_name, gender, pitch, rate)126 127        sentences = split_sentences(text)128        if not sentences:129            sentences = [text.strip()]130 131        all_audio = []132        seed = torch.randint(0, 2**32, (1,)).item()133 134        for sentence in sentences:135            desc_tokens = description_tokenizer(description, return_tensors="pt").to(DEVICE)136            prompt_tokens = tokenizer(sentence, return_tensors="pt").to(DEVICE)137 138            torch.manual_seed(seed)139            if torch.cuda.is_available():140                torch.cuda.manual_seed(seed)141 142            with torch.no_grad():143                generation = model.generate(144                    input_ids=desc_tokens.input_ids,145                    attention_mask=desc_tokens.attention_mask,146                    prompt_input_ids=prompt_tokens.input_ids,147                    prompt_attention_mask=prompt_tokens.attention_mask,148                    do_sample=do_sample,149                    temperature=temperature if do_sample else 1.0,150                    min_new_tokens=10,151                )152 153            audio_chunk = generation.cpu().numpy().squeeze()154            audio_chunk = (audio_chunk * 32767).astype(np.int16)155            all_audio.append(audio_chunk)156 157            # Add 0.3s silence between sentences158            silence = np.zeros(int(SAMPLE_RATE * 0.3), dtype=np.int16)159            all_audio.append(silence)160 161        if not all_audio:162            return None163 164        audio = np.concatenate(all_audio)165        return audio166 167    except Exception as e:168        print(f"Error generating speech: {e}")169        import traceback170        traceback.print_exc()171        return None172 173 174@app.get("/")175async def root():176    """Health check endpoint."""177    return {178        "status": "ok",179        "model": "Indic Parler-TTS",180        "speakers": list(SPEAKERS.keys()),181        "sample_rate": SAMPLE_RATE,182        "endpoints": {183            "POST /tts": "Standard TTS (wait for full audio)",184            "POST /tts/stream": "HTTP Streaming TTS (audio chunks in real-time)",185            "WS /ws/tts": "WebSocket TTS (BEST FOR PIPECAT - true real-time bidirectional)",186            "GET /speakers": "List available speakers"187        },188        "device": DEVICE,189        "optimization": "fp16 (half precision)"190    }191 192 193def generate_audio_chunks_streaming(text, speaker, pitch, rate, temperature, do_sample):194    """Generate audio using official ParlerTTSStreamer for true streaming."""195    text = clean_urdu_text(text)196    speaker_name = SPEAKERS.get(speaker, "")197    gender = "female" if "Female" in speaker or speaker in ["Divya", "Rani"] else "male"198    description = build_description(speaker_name, gender, pitch, rate)199 200    # Create streamer for real-time audio chunks201    play_steps = int(model.config.sampling_rate * 0.5)  # 0.5 second chunks202    streamer = ParlerTTSStreamer(model, device=DEVICE, play_steps=play_steps)203 204    # Tokenize205    desc_tokens = description_tokenizer(description, return_tensors="pt").to(DEVICE)206    prompt_tokens = tokenizer(text, return_tensors="pt").to(DEVICE)207 208    # Set seed209    seed = torch.randint(0, 2**32, (1,)).item()210    torch.manual_seed(seed)211    if torch.cuda.is_available():212        torch.cuda.manual_seed(seed)213 214    # Generate in background thread215    generation_kwargs = dict(216        input_ids=desc_tokens.input_ids,217        attention_mask=desc_tokens.attention_mask,218        prompt_input_ids=prompt_tokens.input_ids,219        prompt_attention_mask=prompt_tokens.attention_mask,220        streamer=streamer,221        do_sample=do_sample,222        temperature=temperature if do_sample else 1.0,223        min_new_tokens=10,224    )225 226    thread = Thread(target=model.generate, kwargs=generation_kwargs)227    thread.daemon = True228    thread.start()229 230    # Yield audio chunks as they're generated231    for audio_chunk in streamer:232        if audio_chunk.shape[0] > 0:233            # ParlerTTSStreamer yields numpy arrays directly, not tensors234            audio_int16 = (audio_chunk * 32767).astype(np.int16)235            yield audio_int16236 237    thread.join()238 239 240async def generate_audio_stream(text, speaker, pitch, rate, temperature, do_sample):241    """Stream audio generation sentence by sentence."""242    text = clean_urdu_text(text)243    speaker_name = SPEAKERS.get(speaker, "")244    gender = "female" if "Female" in speaker or speaker in ["Divya", "Rani"] else "male"245    description = build_description(speaker_name, gender, pitch, rate)246 247    sentences = split_sentences(text)248    if not sentences:249        sentences = [text.strip()]250 251    # Collect all audio chunks252    all_audio = []253    seed = torch.randint(0, 2**32, (1,)).item()254 255    for sentence in sentences:256        desc_tokens = description_tokenizer(description, return_tensors="pt").to(DEVICE)257        prompt_tokens = tokenizer(sentence, return_tensors="pt").to(DEVICE)258 259        torch.manual_seed(seed)260        if torch.cuda.is_available():261            torch.cuda.manual_seed(seed)262 263        with torch.no_grad():264            generation = model.generate(265                input_ids=desc_tokens.input_ids,266                attention_mask=desc_tokens.attention_mask,267                prompt_input_ids=prompt_tokens.input_ids,268                prompt_attention_mask=prompt_tokens.attention_mask,269                do_sample=do_sample,270                temperature=temperature if do_sample else 1.0,271                min_new_tokens=10,272            )273 274        audio_chunk = generation.cpu().numpy().squeeze()275        audio_chunk = (audio_chunk * 32767).astype(np.int16)276        all_audio.append(audio_chunk)277 278        # Add 0.3s silence between sentences279        silence = np.zeros(int(SAMPLE_RATE * 0.3), dtype=np.int16)280        all_audio.append(silence)281 282    # Convert to WAV283    audio = np.concatenate(all_audio)284    audio_buffer = BytesIO()285    sf.write(audio_buffer, audio, SAMPLE_RATE, format='WAV')286    audio_buffer.seek(0)287    return audio_buffer.getvalue()288 289 290@app.post("/tts/stream")291async def text_to_speech_streaming(request: TTSRequest):292    """Generate speech with real-time streaming (fastest latency)."""293    if not request.text.strip():294        raise HTTPException(status_code=400, detail="Text cannot be empty")295 296    if request.speaker not in SPEAKERS:297        raise HTTPException(status_code=400, detail=f"Invalid speaker. Choose from: {list(SPEAKERS.keys())}")298 299    async def audio_generator():300        """Generator that yields audio chunks and WAV header."""301        import struct302 303        try:304            # WAV header will be written first305            wav_header_written = False306 307            for audio_chunk in generate_audio_chunks_streaming(308                request.text,309                request.speaker,310                request.pitch,311                request.rate,312                request.temperature,313                request.do_sample314            ):315                if not wav_header_written:316                    # Write WAV header on first chunk317                    channels = 1318                    sample_width = 2319                    framerate = SAMPLE_RATE320 321                    audio_buffer = BytesIO()322                    sf.write(audio_buffer, audio_chunk, SAMPLE_RATE, format='WAV')323                    audio_buffer.seek(0)324                    wav_data = audio_buffer.read()325 326                    yield wav_data327                    wav_header_written = True328                else:329                    # For subsequent chunks, just append raw audio data330                    yield audio_chunk.tobytes()331 332        except Exception as e:333            print(f"Error in streaming: {e}")334            import traceback335            traceback.print_exc()336 337    return StreamingResponse(338        audio_generator(),339        media_type="audio/wav",340        headers={"Content-Disposition": "inline; filename=speech.wav"}341    )342 343 344@app.post("/tts")345async def text_to_speech(request: TTSRequest):346    """Generate speech from Urdu text."""347    if not request.text.strip():348        raise HTTPException(status_code=400, detail="Text cannot be empty")349 350    if request.speaker not in SPEAKERS:351        raise HTTPException(status_code=400, detail=f"Invalid speaker. Choose from: {list(SPEAKERS.keys())}")352 353    try:354        audio_data = await generate_audio_stream(355            request.text,356            request.speaker,357            request.pitch,358            request.rate,359            request.temperature,360            request.do_sample361        )362 363        if audio_data is None:364            raise HTTPException(status_code=500, detail="Failed to generate speech")365 366        return StreamingResponse(367            iter([audio_data]),368            media_type="audio/wav",369            headers={"Content-Disposition": "attachment; filename=speech.wav"}370        )371    except Exception as e:372        print(f"Error in TTS endpoint: {e}")373        raise HTTPException(status_code=500, detail=str(e))374 375 376@app.get("/speakers")377async def get_speakers():378    """Get list of available speakers."""379    return {"speakers": list(SPEAKERS.keys())}380 381 382@app.websocket("/ws/tts")383async def websocket_tts(websocket: WebSocket):384    """WebSocket endpoint for real-time audio streaming.385 386    Usage:387    1. Connect to ws://localhost:7860/ws/tts388    2. Send JSON: {"text": "سلام دنیا", "speaker": "Divya"}389    3. Receive audio chunks in real-time390    4. Connection closes when generation completes391    """392    await websocket.accept()393    try:394        data = await websocket.receive_text()395        request_data = json.loads(data)396 397        text = request_data.get("text", "").strip()398        speaker = request_data.get("speaker", "Divya")399        pitch = request_data.get("pitch", "Moderate")400        rate = request_data.get("rate", "Moderate")401        temperature = request_data.get("temperature", 0.8)402        do_sample = request_data.get("do_sample", True)403 404        # Validate inputs405        if not text:406            await websocket.send_json({"error": "Text cannot be empty"})407            await websocket.close()408            return409 410        if speaker not in SPEAKERS:411            await websocket.send_json({412                "error": f"Invalid speaker. Choose from: {list(SPEAKERS.keys())}"413            })414            await websocket.close()415            return416 417        # Send status message418        await websocket.send_json({419            "status": "generating",420            "message": f"Generating audio for speaker {speaker}..."421        })422 423        # Generate and stream audio chunks424        chunk_count = 0425        try:426            for audio_chunk in generate_audio_chunks_streaming(427                text, speaker, pitch, rate, temperature, do_sample428            ):429                # Send audio chunk as binary data430                await websocket.send_bytes(audio_chunk.tobytes())431                chunk_count += 1432 433            # Send completion message434            await websocket.send_json({435                "status": "complete",436                "chunks_sent": chunk_count437            })438 439        except Exception as e:440            await websocket.send_json({441                "error": f"Generation failed: {str(e)}"442            })443 444    except WebSocketDisconnect:445        print("WebSocket client disconnected")446    except json.JSONDecodeError:447        await websocket.send_json({"error": "Invalid JSON format"})448        await websocket.close()449    except Exception as e:450        print(f"WebSocket error: {e}")451        await websocket.close()452 453 454if __name__ == "__main__":455    import uvicorn456    uvicorn.run(app, host="0.0.0.0", port=7860)457