211 lines
6.8 KiB
Python
Executable File
211 lines
6.8 KiB
Python
Executable File
#!/usr/bin/env python3
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
import os
|
|
import signal
|
|
import sys
|
|
import time
|
|
import traceback
|
|
import uuid
|
|
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
from pathlib import Path
|
|
from threading import Lock
|
|
from typing import Any
|
|
|
|
|
|
PROJECT_DIR = Path(__file__).resolve().parents[1]
|
|
OUTPUT_DIR = Path(os.environ.get("KSAY_OUTPUT_DIR", PROJECT_DIR / "generated-audio"))
|
|
DEFAULT_HOST = os.environ.get("KSAY_HOST", "127.0.0.1")
|
|
DEFAULT_PORT = int(os.environ.get("KSAY_PORT", "7332"))
|
|
DEFAULT_MODEL = os.environ.get("KSAY_MODEL", "mlx-community/Kokoro-82M-8bit")
|
|
DEFAULT_VOICE = os.environ.get("KSAY_VOICE", "af_heart")
|
|
DEFAULT_LANG_CODE = os.environ.get("KSAY_LANG_CODE", "a")
|
|
MAX_TEXT_CHARS = int(os.environ.get("KSAY_MAX_TEXT_CHARS", "8000"))
|
|
|
|
|
|
class KokoroEngine:
|
|
def __init__(self, model_name: str):
|
|
self.model_name = model_name
|
|
self.model = None
|
|
self.loaded_at = None
|
|
self.load_seconds = None
|
|
self.lock = Lock()
|
|
|
|
def load(self) -> None:
|
|
started = time.perf_counter()
|
|
from mlx_audio.tts.utils import load_model
|
|
|
|
self.model = load_model(self.model_name)
|
|
self.loaded_at = time.time()
|
|
self.load_seconds = time.perf_counter() - started
|
|
|
|
def synthesize(
|
|
self,
|
|
*,
|
|
text: str,
|
|
voice: str,
|
|
speed: float,
|
|
lang_code: str,
|
|
output: str | None,
|
|
) -> dict[str, Any]:
|
|
if self.model is None:
|
|
raise RuntimeError("Kokoro model is not loaded.")
|
|
|
|
clean_text = text.strip()
|
|
if not clean_text:
|
|
raise ValueError("text is required.")
|
|
clean_text = clean_text[:MAX_TEXT_CHARS]
|
|
|
|
output_path = resolve_output_path(output)
|
|
started = time.perf_counter()
|
|
|
|
with self.lock:
|
|
audio_chunks = []
|
|
sample_rate = None
|
|
for result in self.model.generate(
|
|
text=clean_text,
|
|
voice=voice,
|
|
speed=speed,
|
|
lang_code=lang_code,
|
|
):
|
|
audio_chunks.append(result.audio)
|
|
sample_rate = result.sample_rate
|
|
|
|
if not audio_chunks:
|
|
raise RuntimeError("Kokoro did not return audio.")
|
|
|
|
import mlx.core as mx
|
|
import numpy as np
|
|
from mlx_audio.audio_io import write as audio_write
|
|
|
|
audio = (
|
|
mx.concatenate(audio_chunks, axis=0)
|
|
if len(audio_chunks) > 1
|
|
else audio_chunks[0]
|
|
)
|
|
audio_write(str(output_path), np.array(audio), sample_rate, format="wav")
|
|
|
|
elapsed = time.perf_counter() - started
|
|
return {
|
|
"ok": True,
|
|
"filePath": str(output_path),
|
|
"model": self.model_name,
|
|
"voice": voice,
|
|
"speed": speed,
|
|
"langCode": lang_code,
|
|
"sampleRate": sample_rate,
|
|
"segments": len(audio_chunks),
|
|
"seconds": round(elapsed, 3),
|
|
"characters": len(clean_text),
|
|
}
|
|
|
|
|
|
def resolve_output_path(output: str | None) -> Path:
|
|
if output:
|
|
path = Path(output).expanduser()
|
|
if path.suffix.lower() != ".wav":
|
|
path = path.with_suffix(".wav")
|
|
if not path.is_absolute():
|
|
path = (Path.cwd() / path).resolve()
|
|
else:
|
|
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
|
path = OUTPUT_DIR / f"ksay-{int(time.time() * 1000)}-{uuid.uuid4().hex[:6]}.wav"
|
|
|
|
path.parent.mkdir(parents=True, exist_ok=True)
|
|
return path
|
|
|
|
|
|
def parse_json_body(handler: BaseHTTPRequestHandler) -> dict[str, Any]:
|
|
length = int(handler.headers.get("content-length", "0"))
|
|
if length <= 0:
|
|
return {}
|
|
body = handler.rfile.read(length)
|
|
return json.loads(body.decode("utf-8"))
|
|
|
|
|
|
def make_handler(engine: KokoroEngine):
|
|
class KsayHandler(BaseHTTPRequestHandler):
|
|
server_version = "ksay-kokoro/0.1"
|
|
|
|
def log_message(self, fmt: str, *args: Any) -> None:
|
|
sys.stderr.write("%s - %s\n" % (self.log_date_time_string(), fmt % args))
|
|
|
|
def write_json(self, status: int, value: dict[str, Any]) -> None:
|
|
body = json.dumps(value).encode("utf-8")
|
|
self.send_response(status)
|
|
self.send_header("content-type", "application/json")
|
|
self.send_header("content-length", str(len(body)))
|
|
self.end_headers()
|
|
self.wfile.write(body)
|
|
|
|
def do_GET(self) -> None:
|
|
if self.path == "/health":
|
|
self.write_json(
|
|
200,
|
|
{
|
|
"ok": True,
|
|
"model": engine.model_name,
|
|
"loaded": engine.model is not None,
|
|
"loadedAt": engine.loaded_at,
|
|
"loadSeconds": engine.load_seconds,
|
|
"defaultVoice": DEFAULT_VOICE,
|
|
"defaultLangCode": DEFAULT_LANG_CODE,
|
|
},
|
|
)
|
|
return
|
|
self.write_json(404, {"ok": False, "error": "not found"})
|
|
|
|
def do_POST(self) -> None:
|
|
if self.path != "/say":
|
|
self.write_json(404, {"ok": False, "error": "not found"})
|
|
return
|
|
|
|
try:
|
|
payload = parse_json_body(self)
|
|
result = engine.synthesize(
|
|
text=str(payload.get("text", "")),
|
|
voice=str(payload.get("voice") or DEFAULT_VOICE),
|
|
speed=float(payload.get("speed") or 1.0),
|
|
lang_code=str(payload.get("langCode") or DEFAULT_LANG_CODE),
|
|
output=payload.get("output"),
|
|
)
|
|
self.write_json(200, result)
|
|
except Exception as exc:
|
|
traceback.print_exc()
|
|
self.write_json(500, {"ok": False, "error": str(exc)})
|
|
|
|
return KsayHandler
|
|
|
|
|
|
def main() -> int:
|
|
parser = argparse.ArgumentParser(description="Warm Kokoro TTS server for ksay.")
|
|
parser.add_argument("--host", default=DEFAULT_HOST)
|
|
parser.add_argument("--port", type=int, default=DEFAULT_PORT)
|
|
parser.add_argument("--model", default=DEFAULT_MODEL)
|
|
args = parser.parse_args()
|
|
|
|
engine = KokoroEngine(args.model)
|
|
print(f"Loading {args.model}...", flush=True)
|
|
engine.load()
|
|
print(
|
|
f"ksay Kokoro ready on http://{args.host}:{args.port} "
|
|
f"after {engine.load_seconds:.2f}s",
|
|
flush=True,
|
|
)
|
|
|
|
httpd = ThreadingHTTPServer((args.host, args.port), make_handler(engine))
|
|
|
|
def shutdown(_signum: int, _frame: Any) -> None:
|
|
httpd.shutdown()
|
|
|
|
signal.signal(signal.SIGTERM, shutdown)
|
|
signal.signal(signal.SIGINT, shutdown)
|
|
httpd.serve_forever()
|
|
return 0
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|