bilalRHCH commited on
Commit
c9523f7
·
verified ·
1 Parent(s): 8233cad

Upload server.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. server.py +267 -0
server.py ADDED
@@ -0,0 +1,267 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import logging
2
+ import os
3
+ import io
4
+ import numpy as np
5
+ import torch
6
+ import soundfile as sf
7
+ import uuid
8
+ import time
9
+ from fastapi import FastAPI, HTTPException, BackgroundTasks
10
+ from fastapi.responses import StreamingResponse, JSONResponse
11
+ from fastapi.staticfiles import StaticFiles
12
+ from fastapi.middleware.cors import CORSMiddleware
13
+ from pydantic import BaseModel
14
+ from typing import Optional
15
+
16
+ from omnivoice import OmniVoice, OmniVoiceGenerationConfig
17
+ from text_preprocessor import chunk_text
18
+
19
+ logging.basicConfig(
20
+ level=logging.WARNING,
21
+ format="%(asctime)s %(name)s %(levelname)s: %(message)s",
22
+ )
23
+ logger = logging.getLogger(__name__)
24
+
25
+ # FastAPI app
26
+ app = FastAPI(title="Arabic TTS Server (OmniVoice)")
27
+
28
+ app.add_middleware(
29
+ CORSMiddleware,
30
+ allow_origins=[
31
+ "https://arabic-tts-frontend.web.app",
32
+ "https://arabic-tts-frontend.firebaseapp.com",
33
+ "http://localhost:3000",
34
+ "http://localhost:8000"
35
+ ],
36
+ allow_credentials=True,
37
+ allow_methods=["*"],
38
+ allow_headers=["*"],
39
+ )
40
+
41
+ # Global variables for model
42
+ CHECKPOINT = os.environ.get("OMNIVOICE_MODEL", "k2-fsa/OmniVoice")
43
+ model = None
44
+ sampling_rate = 24000
45
+
46
+ # Simple In-Memory Database for State Tracking (Step 4 preview)
47
+ tasks_db = {}
48
+
49
+ @app.on_event("startup")
50
+ async def startup_event():
51
+ global model, sampling_rate
52
+ print(f"Loading OmniVoice model from {CHECKPOINT} ...")
53
+ model = OmniVoice.from_pretrained(
54
+ CHECKPOINT,
55
+ load_asr=True,
56
+ device_map="cpu", # Using CPU by default, adjust if a GPU is available.
57
+ )
58
+ sampling_rate = model.sampling_rate
59
+ print("Model loaded successfully!")
60
+
61
+ class SynthesizeRequest(BaseModel):
62
+ text: str
63
+ voice: Optional[str] = "Auto"
64
+ speed: Optional[float] = 1.0
65
+
66
+ import os
67
+ import shutil
68
+
69
+ def process_audio_task(task_id: str, chunks: list[str], speed: float, voice_id: str):
70
+ """
71
+ Background worker that iterivately generates audio, saves chunks to disk,
72
+ concatenates them, and handles cleanup. Validates consistent voice!
73
+ """
74
+ try:
75
+ total_chunks = len(chunks)
76
+ tasks_db[task_id]["status"] = "processing"
77
+ tasks_db[task_id]["total_chunks"] = total_chunks
78
+
79
+ chunk_dir = os.path.join("audio_chunks", task_id)
80
+ os.makedirs(chunk_dir, exist_ok=True)
81
+
82
+ gen_config = OmniVoiceGenerationConfig(
83
+ num_step=32,
84
+ guidance_scale=2.0,
85
+ denoise=True,
86
+ preprocess_prompt=False,
87
+ postprocess_output=False,
88
+ )
89
+
90
+ master_voice_prompt = None
91
+
92
+ # Check if user requested a specific built-in voice
93
+ if voice_id and voice_id != "Auto":
94
+ voice_path = os.path.join("voices", f"{voice_id}.wav")
95
+ text_path = os.path.join("voices", f"{voice_id}.txt")
96
+ if os.path.exists(voice_path):
97
+ ref_text = None
98
+ if os.path.exists(text_path):
99
+ with open(text_path, "r", encoding="utf-8") as f:
100
+ ref_text = f.read().strip()
101
+ try:
102
+ master_voice_prompt = model.create_voice_clone_prompt(ref_audio=voice_path, ref_text=ref_text)
103
+ except Exception as e:
104
+ logger.warning(f"Voice clone setup failed: {e}")
105
+
106
+ for i, chunk in enumerate(chunks):
107
+ # Update state tracker
108
+ tasks_db[task_id]["current_chunk"] = i + 1
109
+ tasks_db[task_id]["progress"] = int(((i) / total_chunks) * 100)
110
+
111
+ chunk_path = os.path.join(chunk_dir, f"chunk_{i}.wav")
112
+
113
+ # Check if already generated for resume ability
114
+ if not os.path.exists(chunk_path):
115
+ kw = dict(
116
+ text=chunk,
117
+ language="Auto",
118
+ generation_config=gen_config
119
+ )
120
+ if speed is not None and speed != 1.0:
121
+ kw["speed"] = speed
122
+
123
+ # Apply consistent voice cloning (prevents mid-book gender switching)
124
+ if master_voice_prompt is not None:
125
+ kw["voice_clone_prompt"] = master_voice_prompt
126
+
127
+ # Generate Audio via OmniVoice
128
+ audio = model.generate(**kw)
129
+ waveform = audio[0].squeeze(0).numpy()
130
+
131
+ # Save chunk to disk incrementally
132
+ sf.write(chunk_path, waveform, sampling_rate, format='wav', subtype='PCM_16')
133
+
134
+ # If Auto mode, use the FIRST successfully generated chunk as the reference
135
+ # voice for ALL subsequent chunks. This locks the randomly chosen voice!
136
+ if master_voice_prompt is None and i == 0:
137
+ try:
138
+ master_voice_prompt = model.create_voice_clone_prompt(ref_audio=chunk_path, ref_text=chunk)
139
+ except Exception as e:
140
+ logger.warning(f"Could not extract voice clone from chunk 0: {e}")
141
+
142
+ # All chunks generated, now concatenate
143
+ tasks_db[task_id]["status"] = "stitching"
144
+
145
+ all_data = []
146
+ sr = sampling_rate
147
+ for i in range(total_chunks):
148
+ chunk_path = os.path.join(chunk_dir, f"chunk_{i}.wav")
149
+ data, sr = sf.read(chunk_path)
150
+ all_data.append(data)
151
+
152
+ final_waveform = np.concatenate(all_data)
153
+
154
+ final_dir = os.path.join("static", "audio")
155
+ os.makedirs(final_dir, exist_ok=True)
156
+ final_path = os.path.join(final_dir, f"{task_id}.wav")
157
+
158
+ sf.write(final_path, final_waveform, sr, format='wav', subtype='PCM_16')
159
+
160
+ # Cleanup temporary chunks
161
+ shutil.rmtree(chunk_dir)
162
+
163
+ tasks_db[task_id]["status"] = "completed"
164
+ tasks_db[task_id]["progress"] = 100
165
+ tasks_db[task_id]["download_url"] = f"audio/{task_id}.wav"
166
+
167
+ except Exception as e:
168
+ logger.error(f"Background task failed: {str(e)}")
169
+ tasks_db[task_id]["status"] = "failed"
170
+ tasks_db[task_id]["error"] = str(e)
171
+
172
+ @app.post("/synthesize")
173
+ async def synthesize(req: SynthesizeRequest, background_tasks: BackgroundTasks):
174
+ if not model:
175
+ raise HTTPException(status_code=500, detail="Model not loaded yet.")
176
+
177
+ try:
178
+ # Step 1 Integration: chunk the text
179
+ chunks = chunk_text(req.text.strip())
180
+ if not chunks:
181
+ raise HTTPException(status_code=400, detail="Text is empty or invalid.")
182
+
183
+ task_id = str(uuid.uuid4())
184
+
185
+ # Step 4 preview: Initialize state tracking
186
+ tasks_db[task_id] = {
187
+ "task_id": task_id,
188
+ "status": "pending",
189
+ "progress": 0,
190
+ "current_chunk": 0,
191
+ "total_chunks": len(chunks)
192
+ }
193
+
194
+ # Step 2 Integration: Start the background process instead of blocking
195
+ background_tasks.add_task(process_audio_task, task_id, chunks, req.speed, req.voice)
196
+
197
+ return JSONResponse(content={"task_id": task_id, "message": "Audio generation started in the background."})
198
+
199
+ except Exception as e:
200
+ logger.error(f"Synthesis failed: {str(e)}")
201
+ raise HTTPException(status_code=500, detail=str(e))
202
+
203
+ @app.get("/status/{task_id}")
204
+ async def get_status(task_id: str):
205
+ if task_id not in tasks_db:
206
+ raise HTTPException(status_code=404, detail="Task not found.")
207
+ return JSONResponse(content=tasks_db[task_id])
208
+ # ── Distributed Worker: /synthesize_chunk endpoint ───────────────────────────
209
+ class ChunkRequest(BaseModel):
210
+ text: str
211
+ voice: Optional[str] = "Auto"
212
+ speed: Optional[float] = 1.0
213
+ chunk_index: int = 0
214
+
215
+ @app.post("/synthesize_chunk")
216
+ async def synthesize_chunk(req: ChunkRequest):
217
+ """
218
+ Lightweight endpoint for distributed workers.
219
+ Receives a single text chunk, generates audio, and returns it as bytes.
220
+ The Orchestrator calls this on multiple Workers simultaneously.
221
+ """
222
+ if not model:
223
+ raise HTTPException(status_code=500, detail="Model not loaded.")
224
+
225
+ gen_config = OmniVoiceGenerationConfig(
226
+ num_step=32,
227
+ guidance_scale=2.0,
228
+ denoise=True,
229
+ preprocess_prompt=False,
230
+ postprocess_output=False,
231
+ )
232
+
233
+ kw = dict(text=req.text.strip(), language="Auto", generation_config=gen_config)
234
+ if req.speed != 1.0:
235
+ kw["speed"] = req.speed
236
+
237
+ # Voice selection
238
+ if req.voice and req.voice != "Auto":
239
+ voice_path = os.path.join("voices", f"{req.voice}.wav")
240
+ if os.path.exists(voice_path):
241
+ try:
242
+ kw["voice_clone_prompt"] = model.create_voice_clone_prompt(ref_audio=voice_path)
243
+ except Exception as e:
244
+ logger.warning(f"Voice clone failed: {e}")
245
+
246
+ try:
247
+ audio = model.generate(**kw)
248
+ waveform = audio[0].squeeze(0).numpy()
249
+
250
+ # Return the waveform as bytes in a WAV container
251
+ buffer = io.BytesIO()
252
+ sf.write(buffer, waveform, sampling_rate, format='wav', subtype='PCM_16')
253
+ buffer.seek(0)
254
+
255
+ headers = {"X-Chunk-Index": str(req.chunk_index)}
256
+ return StreamingResponse(buffer, media_type="audio/wav", headers=headers)
257
+
258
+ except Exception as e:
259
+ raise HTTPException(status_code=500, detail=str(e))
260
+
261
+ # Ensure voices directory is explicitly available
262
+ os.makedirs("voices", exist_ok=True)
263
+ app.mount("/voices", StaticFiles(directory="voices"), name="voices")
264
+
265
+ # Mount static files directly on root
266
+ app.mount("/", StaticFiles(directory="static", html=True), name="static")
267
+