Repository navigation
Expand file tree
/
Copy pathkokoro_tts.py
More file actions
135 lines (115 loc) · 4.45 KB
/
Copy pathkokoro_tts.py
File metadata and controls
135 lines (115 loc) · 4.45 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
"""
Kokoro-82M TTS wrapper.
Mirrors the Piper loader pattern: lazily downloads model files to voices/kokoro/
on first use, loads a single Kokoro instance, and exposes synth(text, voice)
returning (float32 numpy array, sample_rate).
The kokoro-onnx package ships two file paths we have to provide explicitly:
- an ONNX model file (~310 MB)
- a voices .bin file (~6 MB, contains every speaker embedding)
Usage:
k = KokoroTTS(Path("voices"))
audio, sr = k.synth("hello world", voice="af_heart")
"""
from __future__ import annotations
import urllib.request
import warnings
from pathlib import Path
from threading import Lock
import numpy as np
# Silence the cosmetic RequestsDependencyWarning that the `requests` library
# emits on import when urllib3/chardet versions don't match its metadata —
# triggered indirectly when kokoro-onnx pulls in huggingface-hub.
warnings.filterwarnings(
"ignore",
category=UserWarning,
module=r"requests(\..*)?",
)
_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"
)
# File size floors used to detect partial/corrupt downloads.
_MIN_MODEL_BYTES = 100_000_000 # ~310 MB real; floor well below that
_MIN_VOICES_BYTES = 1_000_000 # ~6 MB real
def _download_with_progress(url: str, dest: Path) -> None:
"""Download via urlretrieve to a .tmp file, atomic-rename to dest."""
tmp = dest.with_name(dest.name + ".tmp")
last_pct = [-1]
def hook(block_num: int, block_size: int, total_size: int) -> None:
if total_size <= 0:
return
pct = min(100, block_num * block_size * 100 // total_size)
if pct == last_pct[0]:
return
last_pct[0] = pct
mb = block_num * block_size / (1024 * 1024)
total_mb = total_size / (1024 * 1024)
print(
f"\r[kokoro] downloading {dest.name}: {pct:3d}% "
f"({mb:.1f} / {total_mb:.1f} MB)",
end="",
flush=True,
)
try:
urllib.request.urlretrieve(url, tmp, reporthook=hook)
print()
tmp.replace(dest)
except Exception:
if tmp.exists():
tmp.unlink(missing_ok=True)
raise
def _ensure_file(dest: Path, url: str, min_bytes: int) -> None:
"""Ensure `dest` exists and is at least `min_bytes`. Downloads if not."""
if dest.exists():
try:
size = dest.stat().st_size
except OSError:
size = 0
if size >= min_bytes:
return
# Too small — treat as corrupt, redownload.
print(
f"[kokoro] '{dest.name}' looks truncated ({size} bytes), "
f"deleting and redownloading"
)
dest.unlink(missing_ok=True)
dest.parent.mkdir(parents=True, exist_ok=True)
_download_with_progress(url, dest)
class KokoroTTS:
def __init__(self, voices_dir: Path) -> None:
self._dir = voices_dir / "kokoro"
self._dir.mkdir(parents=True, exist_ok=True)
self._model_path = self._dir / "kokoro-v1.0.onnx"
self._voices_path = self._dir / "voices-v1.0.bin"
self._kokoro = None
self._load_lock = Lock()
def _ensure_loaded(self) -> None:
if self._kokoro is not None:
return
with self._load_lock:
if self._kokoro is not None:
return
_ensure_file(self._model_path, _MODEL_URL, _MIN_MODEL_BYTES)
_ensure_file(self._voices_path, _VOICES_URL, _MIN_VOICES_BYTES)
print("[kokoro] loading model...", flush=True)
from kokoro_onnx import Kokoro
self._kokoro = Kokoro(str(self._model_path), str(self._voices_path))
print("[kokoro] ready", flush=True)
def synth(
self, text: str, voice: str, speed: float = 1.0
) -> tuple[np.ndarray, int]:
"""Synthesize text. Returns (audio float32, sample_rate).
`speed` is a proper time-stretch (pitch-preserved): 1.0 = normal,
1.5 = 50% faster, 0.8 = 20% slower."""
self._ensure_loaded()
if self._kokoro is None:
raise RuntimeError("KokoroTTS.synth called before _ensure_loaded")
audio, sample_rate = self._kokoro.create(
text, voice=voice, speed=float(speed), lang="en-us"
)
audio = np.asarray(audio, dtype=np.float32)
return audio, int(sample_rate)