Jacky2305's picture
Update main.py for Kokoro
d566687
Raw History Blame Contribute Delete
3.05 kB
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel
from kokoro_onnx import Kokoro
import soundfile as sf
import uuid
from fastapi.responses import Response
from io import BytesIO
from pathlib import Path
import requests
app = FastAPI(title="Kokoro 82M TTS API")
MODEL_DIR = Path("/tmp/kokoro_models")
MODEL_DIR.mkdir(exist_ok=True)
MODEL_PATH = MODEL_DIR / "kokoro-v1.0.onnx"
VOICES_PATH = MODEL_DIR / "voices-v1.0.bin"
MODEL_URL = "https://github.com/thewh1teagle/kokoro-onnx/releases/download/model-files-v1.0/kokoro-v1.0.onnx"
VOICES_URL = "https://github.com/thewh1teagle/kokoro-onnx/releases/download/model-files-v1.0/voices-v1.0.bin"
kokoro: Kokoro | None = None
def download_file(url: str, dest: Path):
if dest.exists():
return
print(f"Downloading {dest.name}...")
with requests.get(url, stream=True, timeout=120, allow_redirects=True) as r:
r.raise_for_status()
with open(dest, "wb") as f:
for chunk in r.iter_content(chunk_size=8192):
f.write(chunk)
print(f"✓ Downloaded {dest.name}")
def get_kokoro() -> Kokoro:
global kokoro
if kokoro is not None:
return kokoro
download_file(MODEL_URL, MODEL_PATH)
download_file(VOICES_URL, VOICES_PATH)
kokoro = Kokoro(str(MODEL_PATH), str(VOICES_PATH))
print(f"✓ Kokoro loaded. Voices: {len(kokoro.get_voices())}")
return kokoro
@app.on_event("startup")
async def warmup():
try:
get_kokoro()
print("Warmup complete")
except Exception as e:
print(f"Warmup failed: {e}")
class TTSRequest(BaseModel):
text: str
voice: str = "af_heart"
speed: float = 1.0
sample_rate: int = 24000
@app.post("/tts")
async def tts(request: TTSRequest):
k = get_kokoro()
available = k.get_voices()
if request.voice not in available:
raise HTTPException(400, f"Voice '{request.voice}' not available. Total: {len(available)}")
try:
audio, sr = k.create(request.text, voice=request.voice, speed=request.speed, lang="en-us")
buffer = BytesIO()
sf.write(buffer, audio, sr, format="WAV", subtype="PCM_16")
buffer.seek(0)
return Response(
content=buffer.read(),
media_type="audio/wav",
headers={"Content-Disposition": f"attachment; filename=tts_{uuid.uuid4().hex}.wav"},
)
except Exception as e:
raise HTTPException(500, f"Synthesis failed: {str(e)}")
@app.get("/health")
async def health():
k = get_kokoro()
return {"status": "healthy", "model": "kokoro-v1.0", "sample_rate": 24000, "voices_count": len(k.get_voices()), "voices_preview": k.get_voices()[:10]}
@app.get("/")
async def root():
k = get_kokoro()
return {
"message": "Kokoro 82M TTS API",
"model": "kokoro-v1.0 (82M, Apache 2.0)",
"sample_rate": 24000,
"endpoint": "/tts (POST)",
"voices_total": len(k.get_voices()),
"example": {"text": "Hello!", "voice": "af_heart", "speed": 1.0},
}