Jacky2305's picture
Switch to Gradio SDK
9a9e06c
Raw History Blame Contribute Delete
4.64 kB
import gradio as gr
from kokoro_onnx import Kokoro
import soundfile as sf
import numpy as np
from pathlib import Path
import requests
import tempfile
import os
# Model setup
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 = None
def download_file(url, dest):
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():
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
# Warm up
try:
get_kokoro()
except Exception as e:
print(f"Warmup failed: {e}")
def tts(text, voice, speed):
if not text.strip():
return None, "请输入文本"
k = get_kokoro()
voices = k.get_voices()
if voice not in voices:
return None, f"语音 '{voice}' 不可用"
try:
audio, sr = k.create(text, voice=voice, speed=speed, lang="en-us")
# Save to temp file for Gradio
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
sf.write(f.name, audio, sr, format="WAV", subtype="PCM_16")
return f.name, "生成成功"
except Exception as e:
return None, f"生成失败: {e}"
# Get voice list
try:
voices_list = get_kokoro().get_voices()
except:
voices_list = ["af_heart", "af_bella", "am_michael", "am_adam", "bf_emma", "bm_george"]
# Categorize voices
us_female = [v for v in voices_list if v.startswith("af_")]
us_male = [v for v in voices_list if v.startswith("am_")]
uk_female = [v for v in voices_list if v.startswith("bf_")]
uk_male = [v for v in voices_list if v.startswith("bm_")]
other = [v for v in voices_list if not v.startswith(("af_", "am_", "bf_", "bm_"))]
voice_choices = us_female + us_male + uk_female + uk_male + other
with gr.Blocks(title="Kokoro 82M TTS") as demo:
gr.Markdown("# 🗣️ Kokoro 82M TTS - High Quality CPU TTS")
gr.Markdown("**Model**: Kokoro 82M (Apache 2.0) | **Quality**: UTMOS MOS ~4.45 | **Speed**: RTF ~0.64 on CPU")
with gr.Row():
with gr.Column(scale=2):
text = gr.Textbox(
label="Text to synthesize",
placeholder="Enter text here...",
lines=3,
value="Hello, this is Kokoro 82M text to speech!"
)
voice = gr.Dropdown(
label="Voice",
choices=voice_choices,
value="af_heart" if "af_heart" in voice_choices else voice_choices[0],
info=f"Total {len(voice_choices)} voices available"
)
speed = gr.Slider(0.5, 2.0, value=1.0, step=0.1, label="Speed")
btn = gr.Button("Generate", variant="primary")
status = gr.Textbox(label="Status", interactive=False)
with gr.Column(scale=1):
audio_out = gr.Audio(label="Output", type="filepath")
btn.click(tts, inputs=[text, voice, speed], outputs=[audio_out, status])
with gr.Accordion("Voice Categories", open=False):
gr.Markdown(f"""
**US Female** ({len(us_female)}): {', '.join(us_female[:10])}{'...' if len(us_female) > 10 else ''}
**US Male** ({len(us_male)}): {', '.join(us_male[:10])}{'...' if len(us_male) > 10 else ''}
**UK Female** ({len(uk_female)}): {', '.join(uk_female[:10])}{'...' if len(uk_female) > 10 else ''}
**UK Male** ({len(uk_male)}): {', '.join(uk_male[:10])}{'...' if len(uk_male) > 10 else ''}
**Other** ({len(other)}): {', '.join(other[:10])}{'...' if len(other) > 10 else ''}
""")
gr.Markdown("""
---
**API**: Also available via REST:
- `POST /tts` - `{"text": "...", "voice": "af_heart", "speed": 1.0}` → WAV
- `GET /health` - Health check
- `GET /` - Voice catalog
""")
if __name__ == "__main__":
demo.launch(server_name="0.0.0.0", server_port=7860)