Bluestrikeai commited on
Commit
44a78af
·
verified ·
1 Parent(s): 5159bf6

Create main.py

Browse files
Files changed (1) hide show
  1. main.py +75 -0
main.py ADDED
@@ -0,0 +1,75 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import re
3
+ import io
4
+ import numpy as np
5
+ import scipy.io.wavfile as wavfile
6
+ from fastapi import FastAPI, Depends, HTTPException, Security
7
+ from fastapi.security.api_key import APIKeyHeader
8
+ from fastapi.responses import Response
9
+ from pydantic import BaseModel
10
+
11
+ # Kokoro imports
12
+ from kokoro import KPipeline
13
+
14
+ app = FastAPI(title="Kokoro-82M TTS Backend")
15
+
16
+ # Security setup
17
+ API_KEY_NAME = "X-API-KEY"
18
+ api_key_header = APIKeyHeader(name=API_KEY_NAME, auto_error=False)
19
+
20
+ def get_api_key(api_key: str = Security(api_key_header)):
21
+ expected_key = os.environ.get("SECRET_API_KEY")
22
+ if not expected_key or api_key != expected_key:
23
+ raise HTTPException(status_code=403, detail="Invalid or missing API Key")
24
+ return api_key
25
+
26
+ # Singleton Model Loading
27
+ # Loads 'a' for American English. Change to 'b' for British.
28
+ pipeline = None
29
+
30
+ @app.on_event("startup")
31
+ def load_model():
32
+ global pipeline
33
+ # Initialize the model once into memory
34
+ pipeline = KPipeline(lang_code='a')
35
+
36
+ class TTSRequest(BaseModel):
37
+ text: str
38
+ voice_id: str = "af_bella"
39
+
40
+ def chunk_text(text: str):
41
+ """Splits text by punctuation to avoid Kokoro's context limit."""
42
+ # Split by ., !, ? followed by a space
43
+ chunks = re.split(r'(?<=[.!?]) +', text.strip())
44
+ return [c for c in chunks if c.strip()]
45
+
46
+ @app.post("/generate", dependencies=[Depends(get_api_key)])
47
+ async def generate_audio(request: TTSRequest):
48
+ try:
49
+ chunks = chunk_text(request.text)
50
+ audio_pieces = []
51
+ sample_rate = 24000 # Kokoro default
52
+
53
+ # Process chunks sequentially
54
+ for chunk in chunks:
55
+ # pipeline returns a generator of (graphemes, phonemes, audio)
56
+ generator = pipeline(chunk, voice=request.voice_id, speed=1.0, split_pattern=None)
57
+ for _, _, audio in generator:
58
+ if audio is not None:
59
+ audio_pieces.append(audio)
60
+
61
+ if not audio_pieces:
62
+ raise HTTPException(status_code=400, detail="Could not generate audio from input text.")
63
+
64
+ # Concatenate all numpy audio chunks
65
+ final_audio = np.concatenate(audio_pieces)
66
+
67
+ # Convert to standard WAV bytes in memory
68
+ wav_io = io.BytesIO()
69
+ wavfile.write(wav_io, sample_rate, final_audio)
70
+ wav_io.seek(0)
71
+
72
+ return Response(content=wav_io.read(), media_type="audio/wav")
73
+
74
+ except Exception as e:
75
+ raise HTTPException(status_code=500, detail=str(e))