ebitlogix/Parler_TTS_API
0
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 