zerotts 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- zerotts/__init__.py +26 -0
- zerotts/audio.py +57 -0
- zerotts/chunking.py +270 -0
- zerotts/cli.py +162 -0
- zerotts/codec.py +154 -0
- zerotts/hub.py +93 -0
- zerotts/synthesizer.py +458 -0
- zerotts/text_norm/__init__.py +11 -0
- zerotts/text_norm/data/LICENSE.soe-vinorm +21 -0
- zerotts/text_norm/data/abbreviations.txt +2258 -0
- zerotts/text_norm/vi_normalizer.py +620 -0
- zerotts/tokenizer.py +139 -0
- zerotts/voices.py +129 -0
- zerotts-0.1.0.dist-info/METADATA +350 -0
- zerotts-0.1.0.dist-info/RECORD +20 -0
- zerotts-0.1.0.dist-info/WHEEL +5 -0
- zerotts-0.1.0.dist-info/entry_points.txt +2 -0
- zerotts-0.1.0.dist-info/licenses/LICENSE +21 -0
- zerotts-0.1.0.dist-info/licenses/NOTICE +36 -0
- zerotts-0.1.0.dist-info/top_level.txt +1 -0
zerotts/__init__.py
ADDED
|
@@ -0,0 +1,26 @@
|
|
|
1
|
+
"""ZeroTTS — Vietnamese zero-shot text-to-speech, ONNX runtime, no PyTorch.
|
|
2
|
+
|
|
3
|
+
from zerotts import ZeroTTS
|
|
4
|
+
|
|
5
|
+
tts = ZeroTTS.from_pretrained("zeroweight-ai/ZeroTTS")
|
|
6
|
+
audio = tts.synthesize("Xin chào các bạn.", voice="maichi")
|
|
7
|
+
tts.save_audio(audio, "out.wav")
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from .hub import DEFAULT_REPO_ID, resolve_model_dir
|
|
11
|
+
from .synthesizer import ZeroTTS
|
|
12
|
+
from .text_norm import normalize_vi_text
|
|
13
|
+
from .voices import Voice, list_voices, load_voice
|
|
14
|
+
|
|
15
|
+
__version__ = "0.1.0"
|
|
16
|
+
|
|
17
|
+
__all__ = [
|
|
18
|
+
"ZeroTTS",
|
|
19
|
+
"normalize_vi_text",
|
|
20
|
+
"Voice",
|
|
21
|
+
"list_voices",
|
|
22
|
+
"load_voice",
|
|
23
|
+
"resolve_model_dir",
|
|
24
|
+
"DEFAULT_REPO_ID",
|
|
25
|
+
"__version__",
|
|
26
|
+
]
|
zerotts/audio.py
ADDED
|
@@ -0,0 +1,57 @@
|
|
|
1
|
+
"""Audio I/O helpers — plain numpy + soundfile/scipy, no torch."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from math import gcd
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
|
|
9
|
+
SAMPLE_RATE = 48_000
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def load_wav_mono(path: str, target_sr: int = SAMPLE_RATE) -> np.ndarray:
|
|
13
|
+
"""Load a wav, downmix to mono, resample to ``target_sr``. Returns (T,) float32."""
|
|
14
|
+
import soundfile as sf
|
|
15
|
+
|
|
16
|
+
wav, sr = sf.read(path, always_2d=True)
|
|
17
|
+
wav = wav.mean(axis=1).astype(np.float32)
|
|
18
|
+
if sr != target_sr:
|
|
19
|
+
from scipy.signal import resample_poly
|
|
20
|
+
|
|
21
|
+
g = gcd(target_sr, sr)
|
|
22
|
+
wav = resample_poly(wav, target_sr // g, sr // g)
|
|
23
|
+
return wav.astype(np.float32)
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def resample_wav(wav: np.ndarray, sr_from: int, sr_to: int) -> np.ndarray:
|
|
27
|
+
"""Resample a 1-D float32 array, matching load_wav_mono's resampler."""
|
|
28
|
+
if sr_from == sr_to:
|
|
29
|
+
return wav.astype(np.float32)
|
|
30
|
+
from scipy.signal import resample_poly
|
|
31
|
+
|
|
32
|
+
g = gcd(sr_from, sr_to)
|
|
33
|
+
return resample_poly(wav, sr_to // g, sr_from // g).astype(np.float32)
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def save_wav(audio: np.ndarray, path: str, sample_rate: int = SAMPLE_RATE) -> None:
|
|
37
|
+
"""Save a (..., T) float array (leading dims squeezed) as 16-bit PCM wav."""
|
|
38
|
+
import soundfile as sf
|
|
39
|
+
|
|
40
|
+
sf.write(path, np.asarray(audio).squeeze(), sample_rate, subtype="PCM_16")
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def concat_with_silence(chunks: list[np.ndarray], silence_sec: float,
|
|
44
|
+
sample_rate: int = SAMPLE_RATE) -> np.ndarray:
|
|
45
|
+
"""Join (1, T) or (T,) chunks with ``silence_sec`` of silence between them."""
|
|
46
|
+
parts = [np.asarray(c).reshape(-1) for c in chunks if np.asarray(c).size]
|
|
47
|
+
if not parts:
|
|
48
|
+
return np.zeros((1, 0), dtype=np.float32)
|
|
49
|
+
if silence_sec > 0 and len(parts) > 1:
|
|
50
|
+
gap = np.zeros(int(silence_sec * sample_rate), dtype=np.float32)
|
|
51
|
+
joined: list = []
|
|
52
|
+
for i, p in enumerate(parts):
|
|
53
|
+
if i:
|
|
54
|
+
joined.append(gap)
|
|
55
|
+
joined.append(p)
|
|
56
|
+
parts = joined
|
|
57
|
+
return np.concatenate(parts).astype(np.float32)[None, :]
|
zerotts/chunking.py
ADDED
|
@@ -0,0 +1,270 @@
|
|
|
1
|
+
"""Punctuation normalization + sentence segmentation/chunking for long-form
|
|
2
|
+
text-to-speech input.
|
|
3
|
+
|
|
4
|
+
No sentence-boundary-detection library is a project dependency (no nltk/
|
|
5
|
+
spacy/pysbd), so this uses a regex-based splitter, then greedily packs
|
|
6
|
+
sentences into chunks bounded by an estimated speaking duration.
|
|
7
|
+
|
|
8
|
+
Splitting falls back through a hierarchy so no single chunk ever exceeds
|
|
9
|
+
the character budget, however the user punctuates (or doesn't):
|
|
10
|
+
1. sentence-ending punctuation (. ! ? …)
|
|
11
|
+
2. commas, if a "sentence" is still too long
|
|
12
|
+
3. word boundaries, if a comma-separated piece is still too long
|
|
13
|
+
4. raw character slicing, as a last resort (e.g. one giant run-on word)
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
from __future__ import annotations
|
|
17
|
+
|
|
18
|
+
import re
|
|
19
|
+
|
|
20
|
+
# Split after ., !, ?, or … followed by whitespace (or end of string), but
|
|
21
|
+
# not on a run of digits/abbreviation-like single letters (e.g. "3.5" or
|
|
22
|
+
# "Mr.") immediately before the dot — kept intentionally simple/heuristic.
|
|
23
|
+
_SENTENCE_END_RE = re.compile(r"(?<=[.!?…])\s+")
|
|
24
|
+
|
|
25
|
+
# Split after a comma followed by whitespace, keeping the comma with the
|
|
26
|
+
# preceding text (used only when a whole "sentence" is over budget).
|
|
27
|
+
_COMMA_RE = re.compile(r"(?<=,)\s+")
|
|
28
|
+
|
|
29
|
+
# Rough average speaking rate used to size chunks by estimated duration
|
|
30
|
+
# rather than raw character count (TTS speaking rate varies by language,
|
|
31
|
+
# so this is a coarse heuristic, not a precise duration predictor).
|
|
32
|
+
_CHARS_PER_SEC = 15.0
|
|
33
|
+
|
|
34
|
+
# ── punctuation normalization ────────────────────────────────────────────────
|
|
35
|
+
# Written text carries breaks the model was never trained to voice. The BPE
|
|
36
|
+
# vocab has a token for ';', but the training corpus is transcribed speech,
|
|
37
|
+
# where a semicolon is rare — the model has no reliable prosody for it, and at
|
|
38
|
+
# inference it tends to read as a hard stop or as nothing at all. A comma is
|
|
39
|
+
# the break it actually stands for out loud, so it is rewritten to one. ':' is
|
|
40
|
+
# left as-is (kept, not folded to a comma). Newlines are worse: the BPE
|
|
41
|
+
# normalizer collapses all whitespace to a single space
|
|
42
|
+
# (zerotts.tokenizer.normalize_text), so a line break vanishes entirely
|
|
43
|
+
# and two unrelated lines run together as one breath. Turning it into a
|
|
44
|
+
# sentence stop preserves the pause the layout meant — unless the line already
|
|
45
|
+
# ends in punctuation, which already carries it.
|
|
46
|
+
|
|
47
|
+
_PAUSE_PUNCT_RE = re.compile(r"[;]")
|
|
48
|
+
|
|
49
|
+
# A newline (plus any surrounding horizontal whitespace, and any run of blank
|
|
50
|
+
# lines) with what immediately precedes it captured, so the replacement can look
|
|
51
|
+
# at whether that character is already punctuation.
|
|
52
|
+
_NEWLINE_RE = re.compile(r"(.?)[ \t]*(?:\r?\n[ \t]*)+")
|
|
53
|
+
|
|
54
|
+
# Characters that already terminate or pause a clause — a line ending in one of
|
|
55
|
+
# these needs no added '.'.
|
|
56
|
+
_EXISTING_PUNCT = set(".!?…,;:—–-\"')]}")
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def normalize_punctuation(text: str) -> str:
|
|
60
|
+
"""Rewrite written-only punctuation into the breaks the model can voice.
|
|
61
|
+
|
|
62
|
+
';' -> ','
|
|
63
|
+
newline -> '. ' when the line does not already end in punctuation,
|
|
64
|
+
otherwise just a space.
|
|
65
|
+
|
|
66
|
+
Semicolons are converted FIRST, so a line ending in one counts as
|
|
67
|
+
already-punctuated when the newline rule runs (it ends in ',' by then)
|
|
68
|
+
and does not also collect a '.'. ':' is untouched by this step, but is
|
|
69
|
+
already in _EXISTING_PUNCT so a line ending in one is also treated as
|
|
70
|
+
already-punctuated.
|
|
71
|
+
"""
|
|
72
|
+
if not text:
|
|
73
|
+
return text
|
|
74
|
+
|
|
75
|
+
text = _PAUSE_PUNCT_RE.sub(",", text)
|
|
76
|
+
|
|
77
|
+
def _repl(m: re.Match) -> str:
|
|
78
|
+
prev = m.group(1)
|
|
79
|
+
if not prev: # leading newline(s) — nothing before them to punctuate
|
|
80
|
+
return ""
|
|
81
|
+
if prev in _EXISTING_PUNCT:
|
|
82
|
+
return f"{prev} "
|
|
83
|
+
return f"{prev}. "
|
|
84
|
+
|
|
85
|
+
text = _NEWLINE_RE.sub(_repl, text)
|
|
86
|
+
return text.strip()
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def split_sentences(text: str) -> list[str]:
|
|
90
|
+
text = text.strip()
|
|
91
|
+
if not text:
|
|
92
|
+
return []
|
|
93
|
+
# Split paragraphs first so blank lines always force a break, then
|
|
94
|
+
# sentence-split within each paragraph.
|
|
95
|
+
sentences: list[str] = []
|
|
96
|
+
for para in re.split(r"\n\s*\n", text):
|
|
97
|
+
para = para.strip()
|
|
98
|
+
if not para:
|
|
99
|
+
continue
|
|
100
|
+
for sent in _SENTENCE_END_RE.split(para):
|
|
101
|
+
sent = sent.strip()
|
|
102
|
+
if sent:
|
|
103
|
+
sentences.append(sent)
|
|
104
|
+
return sentences
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def _pack_pieces(pieces: list[str], max_chars: int) -> list[str]:
|
|
108
|
+
"""Greedily join `pieces` (each already <= max_chars) with a single
|
|
109
|
+
space, filling each output chunk as close to max_chars as possible
|
|
110
|
+
without exceeding it."""
|
|
111
|
+
chunks: list[str] = []
|
|
112
|
+
current: list[str] = []
|
|
113
|
+
current_len = 0
|
|
114
|
+
for piece in pieces:
|
|
115
|
+
piece_len = len(piece) + (1 if current else 0)
|
|
116
|
+
if current and current_len + piece_len > max_chars:
|
|
117
|
+
chunks.append(" ".join(current))
|
|
118
|
+
current, current_len = [], 0
|
|
119
|
+
piece_len = len(piece)
|
|
120
|
+
current.append(piece)
|
|
121
|
+
current_len += piece_len
|
|
122
|
+
if current:
|
|
123
|
+
chunks.append(" ".join(current))
|
|
124
|
+
return chunks
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
def _split_by_chars(text: str, max_chars: int) -> list[str]:
|
|
128
|
+
return [text[i:i + max_chars] for i in range(0, len(text), max_chars)]
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
def _split_by_words(text: str, max_chars: int) -> list[str]:
|
|
132
|
+
words = text.split()
|
|
133
|
+
pieces = _pack_pieces(words, max_chars)
|
|
134
|
+
# A single word longer than max_chars (e.g. a spammed run of characters
|
|
135
|
+
# with no spaces) still needs a hard character split.
|
|
136
|
+
out: list[str] = []
|
|
137
|
+
for piece in pieces:
|
|
138
|
+
if len(piece) <= max_chars:
|
|
139
|
+
out.append(piece)
|
|
140
|
+
else:
|
|
141
|
+
out.extend(_split_by_chars(piece, max_chars))
|
|
142
|
+
return out
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
def _atomize(text: str, max_chars: int) -> list[str]:
|
|
146
|
+
"""Break `text` down until every returned piece is <= max_chars,
|
|
147
|
+
preferring the least disruptive split available (comma > word >
|
|
148
|
+
character)."""
|
|
149
|
+
text = text.strip()
|
|
150
|
+
if not text:
|
|
151
|
+
return []
|
|
152
|
+
if len(text) <= max_chars:
|
|
153
|
+
return [text]
|
|
154
|
+
|
|
155
|
+
comma_parts = [p.strip() for p in _COMMA_RE.split(text) if p.strip()]
|
|
156
|
+
if len(comma_parts) > 1:
|
|
157
|
+
out: list[str] = []
|
|
158
|
+
for part in comma_parts:
|
|
159
|
+
out.extend(_atomize(part, max_chars))
|
|
160
|
+
return out
|
|
161
|
+
|
|
162
|
+
if " " in text.strip():
|
|
163
|
+
return _split_by_words(text, max_chars)
|
|
164
|
+
|
|
165
|
+
return _split_by_chars(text, max_chars)
|
|
166
|
+
|
|
167
|
+
|
|
168
|
+
def chunk_text(text: str, max_chunk_sec: float = 15.0) -> list[str]:
|
|
169
|
+
"""Split `text` into speakable chunks whose estimated duration stays
|
|
170
|
+
near `max_chunk_sec`, preferring to break on sentence-ending
|
|
171
|
+
punctuation, then commas, then words, then raw characters as a last
|
|
172
|
+
resort — so a single chunk never balloons past the budget no matter how
|
|
173
|
+
the input is (or isn't) punctuated.
|
|
174
|
+
"""
|
|
175
|
+
max_chars = max(1, int(max_chunk_sec * _CHARS_PER_SEC))
|
|
176
|
+
sentences = split_sentences(text)
|
|
177
|
+
|
|
178
|
+
atoms: list[str] = []
|
|
179
|
+
for sent in sentences:
|
|
180
|
+
atoms.extend(_atomize(sent, max_chars))
|
|
181
|
+
|
|
182
|
+
return _pack_pieces(atoms, max_chars)
|
|
183
|
+
|
|
184
|
+
|
|
185
|
+
# ── per-segment punctuation cleanup ──────────────────────────────────────────
|
|
186
|
+
# Chunking (above) can cut a sentence mid-flow, leaving a segment that ends on
|
|
187
|
+
# a comma or nothing at all, and/or a stray '.'/'!'/'?'/other symbol stranded
|
|
188
|
+
# in the middle of what is now a shorter, standalone utterance. Interior
|
|
189
|
+
# punctuation reads as a hard stop (or gets silently dropped by the tokenizer)
|
|
190
|
+
# where only a pause belongs, so every mid-segment punctuation mark other than
|
|
191
|
+
# a hyphen — word-joining ('twenty-one') or otherwise —, a slash, or a dot is
|
|
192
|
+
# downgraded to a comma-pause. The segment's own end is where a real stop belongs: whatever
|
|
193
|
+
# terminal mark it already had is kept if it's one of '.', '!', '?' (that's
|
|
194
|
+
# real, meaningful prosody), and anything else trailing there — comma, dash,
|
|
195
|
+
# stray symbol, nothing — is normalized to a single '.'.
|
|
196
|
+
|
|
197
|
+
# Trailing run of whitespace/punctuation, i.e. everything after the last
|
|
198
|
+
# alphanumeric character in the segment.
|
|
199
|
+
_TRAILING_RE = re.compile(r"[^\w]+$", re.UNICODE)
|
|
200
|
+
|
|
201
|
+
# Any mid-segment character that is not a word char, whitespace, hyphen,
|
|
202
|
+
# slash, dot, colon, question mark, exclamation mark, quote (single or
|
|
203
|
+
# double), percent sign, or comma itself gets folded into a comma-pause.
|
|
204
|
+
_MID_PUNCT_RE = re.compile(r"[^\w\s\-/.,:?@!\"'%]", re.UNICODE)
|
|
205
|
+
|
|
206
|
+
# Collapse runs of commas (possibly separated by spaces) into a single
|
|
207
|
+
# ', ' — also fixes the missing space left when a punctuation mark abutting
|
|
208
|
+
# a word (e.g. the '"' in 'Hello"world') is folded straight into a comma.
|
|
209
|
+
_REPEAT_COMMA_RE = re.compile(r"\s*(?:,\s*)+")
|
|
210
|
+
|
|
211
|
+
_END_PUNCT = {".", "!", "?"}
|
|
212
|
+
|
|
213
|
+
|
|
214
|
+
def clean_segment_punctuation(text: str) -> str:
|
|
215
|
+
"""Normalize a single chunked segment's punctuation for TTS:
|
|
216
|
+
|
|
217
|
+
1. every mid-segment punctuation mark other than
|
|
218
|
+
'-', '/', '.', ':', '?', '!', '"', "'", '%' becomes a ',' pause
|
|
219
|
+
2. terminal mark: kept as-is if it's '.', '!', or '?', otherwise
|
|
220
|
+
whatever trails (comma, dash, stray symbol, nothing) -> '.'
|
|
221
|
+
"""
|
|
222
|
+
text = text.strip()
|
|
223
|
+
if not text:
|
|
224
|
+
return text
|
|
225
|
+
|
|
226
|
+
m = _TRAILING_RE.search(text)
|
|
227
|
+
core = text[:m.start()] if m else text
|
|
228
|
+
trailing = text[m.start():].rstrip() if m else ""
|
|
229
|
+
|
|
230
|
+
end_punct = trailing[-1] if trailing and trailing[-1] in _END_PUNCT else "."
|
|
231
|
+
|
|
232
|
+
core = _MID_PUNCT_RE.sub(",", core)
|
|
233
|
+
core = _REPEAT_COMMA_RE.sub(", ", core).strip(" ,")
|
|
234
|
+
|
|
235
|
+
if not core:
|
|
236
|
+
return ""
|
|
237
|
+
return core + end_punct
|
|
238
|
+
|
|
239
|
+
|
|
240
|
+
# ── sample text file loading ─────────────────────────────────────────────────
|
|
241
|
+
# Used by the Gradio demo's "Sample texts" picker.
|
|
242
|
+
|
|
243
|
+
def load_text_samples(path) -> dict[str, str]:
|
|
244
|
+
"""Parse a "### name" delimited sample-text file (see webui/test_samples.txt's
|
|
245
|
+
header for the format) into {name: text}. Missing file -> {}."""
|
|
246
|
+
from pathlib import Path
|
|
247
|
+
|
|
248
|
+
path = Path(path)
|
|
249
|
+
if not path.is_file():
|
|
250
|
+
return {}
|
|
251
|
+
|
|
252
|
+
samples: dict[str, str] = {}
|
|
253
|
+
name = None
|
|
254
|
+
lines: list[str] = []
|
|
255
|
+
|
|
256
|
+
def flush():
|
|
257
|
+
if name is not None:
|
|
258
|
+
samples[name] = "\n".join(lines).strip()
|
|
259
|
+
|
|
260
|
+
for raw_line in path.read_text(encoding="utf-8").splitlines():
|
|
261
|
+
if raw_line.startswith("### "):
|
|
262
|
+
flush()
|
|
263
|
+
name = raw_line[4:].strip()
|
|
264
|
+
lines = []
|
|
265
|
+
elif raw_line.startswith("#"):
|
|
266
|
+
continue
|
|
267
|
+
else:
|
|
268
|
+
lines.append(raw_line)
|
|
269
|
+
flush()
|
|
270
|
+
return samples
|
zerotts/cli.py
ADDED
|
@@ -0,0 +1,162 @@
|
|
|
1
|
+
"""Command-line interface: ``zerotts say`` / ``zerotts voices`` / ``zerotts bench``."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import argparse
|
|
6
|
+
import sys
|
|
7
|
+
import time
|
|
8
|
+
|
|
9
|
+
import numpy as np
|
|
10
|
+
|
|
11
|
+
from . import hub
|
|
12
|
+
from .audio import concat_with_silence
|
|
13
|
+
from .chunking import chunk_text, clean_segment_punctuation, normalize_punctuation
|
|
14
|
+
from .synthesizer import ZeroTTS
|
|
15
|
+
from .text_norm import normalize_vi_text
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def _add_model_args(p: argparse.ArgumentParser) -> None:
|
|
19
|
+
p.add_argument("--model", default=hub.DEFAULT_REPO_ID,
|
|
20
|
+
help="HF repo id or local model directory.")
|
|
21
|
+
p.add_argument("--revision", default=None, help="HF revision to pin.")
|
|
22
|
+
p.add_argument("--threads", type=int, default=4, help="onnxruntime intra-op threads.")
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def _add_sampling_args(p: argparse.ArgumentParser) -> None:
|
|
26
|
+
p.add_argument("--cfg_scale", type=float, default=1.0,
|
|
27
|
+
help=">1 guides toward the voice, at 2x the per-frame cost.")
|
|
28
|
+
p.add_argument("--audio_temperature", type=float, default=0.8)
|
|
29
|
+
p.add_argument("--audio_topk", type=int, default=25)
|
|
30
|
+
p.add_argument("--audio_topp", type=float, default=0.95)
|
|
31
|
+
p.add_argument("--audio_repetition_penalty", type=float, default=1.2,
|
|
32
|
+
help="1.2 is the benchmarked default; 1.0 raises WER.")
|
|
33
|
+
p.add_argument("--seed", type=int, default=None, help="Seed the sampler.")
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def _sampling_kwargs(a: argparse.Namespace) -> dict:
|
|
37
|
+
return {
|
|
38
|
+
"cfg_scale": a.cfg_scale,
|
|
39
|
+
"audio_temperature": a.audio_temperature,
|
|
40
|
+
"audio_topk": a.audio_topk,
|
|
41
|
+
"audio_topp": a.audio_topp,
|
|
42
|
+
"audio_repetition_penalty": a.audio_repetition_penalty,
|
|
43
|
+
}
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def cmd_say(a: argparse.Namespace) -> int:
|
|
47
|
+
if a.seed is not None:
|
|
48
|
+
np.random.seed(a.seed)
|
|
49
|
+
tts = ZeroTTS.from_pretrained(a.model, revision=a.revision,
|
|
50
|
+
intra_op_num_threads=a.threads)
|
|
51
|
+
|
|
52
|
+
text = a.text if a.text != "-" else sys.stdin.read()
|
|
53
|
+
# Before chunking: an expansion is several times longer than what it
|
|
54
|
+
# replaces, and the chunk budget has to size the text the model receives.
|
|
55
|
+
if not a.no_text_norm:
|
|
56
|
+
text = normalize_vi_text(text)
|
|
57
|
+
segments = [text]
|
|
58
|
+
if a.chunk:
|
|
59
|
+
segments = [clean_segment_punctuation(s)
|
|
60
|
+
for s in chunk_text(normalize_punctuation(text),
|
|
61
|
+
max_chunk_sec=a.max_chunk_sec)]
|
|
62
|
+
segments = [s for s in segments if s]
|
|
63
|
+
|
|
64
|
+
t0 = time.perf_counter()
|
|
65
|
+
chunks = []
|
|
66
|
+
for i, seg in enumerate(segments, 1):
|
|
67
|
+
if len(segments) > 1:
|
|
68
|
+
print(f"[{i}/{len(segments)}] {seg[:70]}{'…' if len(seg) > 70 else ''}",
|
|
69
|
+
file=sys.stderr)
|
|
70
|
+
chunks.append(tts.synthesize(seg, voice=a.voice, **_sampling_kwargs(a)))
|
|
71
|
+
audio = concat_with_silence(chunks, a.gap_sec, tts.sample_rate)
|
|
72
|
+
elapsed = time.perf_counter() - t0
|
|
73
|
+
|
|
74
|
+
tts.save_audio(audio, a.out)
|
|
75
|
+
dur = audio.shape[-1] / tts.sample_rate
|
|
76
|
+
speed = dur / elapsed if elapsed > 0 else float("inf")
|
|
77
|
+
print(f"{a.out} {dur:.2f}s audio in {elapsed:.2f}s ({speed:.1f}x realtime)")
|
|
78
|
+
return 0
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def cmd_voices(a: argparse.Namespace) -> int:
|
|
82
|
+
tts = ZeroTTS.from_pretrained(a.model, revision=a.revision, warmup=False)
|
|
83
|
+
names = tts.list_voices()
|
|
84
|
+
if not names:
|
|
85
|
+
print("No voice packs in this model directory.")
|
|
86
|
+
print("This build cannot create voices from audio — see the README "
|
|
87
|
+
"(voice cloning).")
|
|
88
|
+
return 1
|
|
89
|
+
for name in names:
|
|
90
|
+
v = tts.load_voice(name)
|
|
91
|
+
desc = f" {v.description}" if v.description else ""
|
|
92
|
+
print(f"{name:20s} {v.language:4s} {v.n_voice_queries} queries{desc}")
|
|
93
|
+
return 0
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def cmd_bench(a: argparse.Namespace) -> int:
|
|
97
|
+
if a.seed is not None:
|
|
98
|
+
np.random.seed(a.seed)
|
|
99
|
+
tts = ZeroTTS.from_pretrained(a.model, revision=a.revision,
|
|
100
|
+
intra_op_num_threads=a.threads)
|
|
101
|
+
text = a.text
|
|
102
|
+
timings = []
|
|
103
|
+
for i in range(a.runs):
|
|
104
|
+
timing: dict = {}
|
|
105
|
+
audio = tts.synthesize(text, voice=a.voice, timing=timing,
|
|
106
|
+
**_sampling_kwargs(a))
|
|
107
|
+
dur = audio.shape[-1] / tts.sample_rate
|
|
108
|
+
timings.append((timing, dur))
|
|
109
|
+
print(f"run {i + 1}: {dur:.2f}s audio, {timing['total_time']:.2f}s wall, "
|
|
110
|
+
f"TTFF {timing['time_to_first_frame'] * 1000:.0f}ms, "
|
|
111
|
+
f"{timing['n_frames']} frames, {dur / timing['total_time']:.1f}x realtime")
|
|
112
|
+
|
|
113
|
+
wall = float(np.median([t["total_time"] for t, _ in timings]))
|
|
114
|
+
dur = float(np.median([d for _, d in timings]))
|
|
115
|
+
ttff = float(np.median([t["time_to_first_frame"] for t, _ in timings]))
|
|
116
|
+
print(f"\nmedian over {a.runs}: {dur / wall:.1f}x realtime, TTFF {ttff * 1000:.0f}ms, "
|
|
117
|
+
f"{a.threads} threads")
|
|
118
|
+
return 0
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
def main(argv=None) -> int:
|
|
122
|
+
p = argparse.ArgumentParser(prog="zerotts", description="ZeroTTS command line.")
|
|
123
|
+
sub = p.add_subparsers(dest="cmd", required=True)
|
|
124
|
+
|
|
125
|
+
say = sub.add_parser("say", help="Synthesize text to a wav file.")
|
|
126
|
+
say.add_argument("text", help="Text to speak, or '-' to read stdin.")
|
|
127
|
+
say.add_argument("-o", "--out", default="out.wav")
|
|
128
|
+
say.add_argument("-v", "--voice", default=None,
|
|
129
|
+
help="Voice name. Omit for the model's unconditional voice.")
|
|
130
|
+
say.add_argument("--chunk", action="store_true",
|
|
131
|
+
help="Split long text into segments and join the audio.")
|
|
132
|
+
say.add_argument("--max_chunk_sec", type=float, default=15.0)
|
|
133
|
+
say.add_argument("--gap_sec", type=float, default=0.15,
|
|
134
|
+
help="Silence inserted between chunks.")
|
|
135
|
+
say.add_argument("--no_text_norm", action="store_true",
|
|
136
|
+
help="Skip Vietnamese normalization of dates/times/numbers. "
|
|
137
|
+
"Use for non-Vietnamese text — the expansions are "
|
|
138
|
+
"Vietnamese words.")
|
|
139
|
+
_add_model_args(say)
|
|
140
|
+
_add_sampling_args(say)
|
|
141
|
+
say.set_defaults(func=cmd_say)
|
|
142
|
+
|
|
143
|
+
voices = sub.add_parser("voices", help="List available voices.")
|
|
144
|
+
_add_model_args(voices)
|
|
145
|
+
voices.set_defaults(func=cmd_voices)
|
|
146
|
+
|
|
147
|
+
bench = sub.add_parser("bench", help="Measure realtime factor and TTFF.")
|
|
148
|
+
bench.add_argument(
|
|
149
|
+
"--text",
|
|
150
|
+
default="Xin chào, đây là một bài kiểm tra tốc độ tổng hợp giọng nói.")
|
|
151
|
+
bench.add_argument("-v", "--voice", default=None)
|
|
152
|
+
bench.add_argument("--runs", type=int, default=3)
|
|
153
|
+
_add_model_args(bench)
|
|
154
|
+
_add_sampling_args(bench)
|
|
155
|
+
bench.set_defaults(func=cmd_bench)
|
|
156
|
+
|
|
157
|
+
a = p.parse_args(argv)
|
|
158
|
+
return a.func(a)
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
if __name__ == "__main__":
|
|
162
|
+
raise SystemExit(main())
|
zerotts/codec.py
ADDED
|
@@ -0,0 +1,154 @@
|
|
|
1
|
+
"""onnxruntime-only MOSS-Audio-Tokenizer-Nano **decoder** — codes -> waveform.
|
|
2
|
+
|
|
3
|
+
Vendored, not downloaded. The decoder graphs ship inside the ZeroTTS weights
|
|
4
|
+
repo (``onnx/codec/``) so a ZeroTTS install has no runtime dependency on any
|
|
5
|
+
third-party model repo staying up or unchanged.
|
|
6
|
+
|
|
7
|
+
Decoder only. The upstream export also has an encoder graph (waveform -> codes),
|
|
8
|
+
used solely to turn reference audio into voice latents. That is voice cloning,
|
|
9
|
+
which this release does not do (see zerotts.voices), so the encoder is neither
|
|
10
|
+
shipped nor wrapped — it would be ~45MB of weights nothing here can call.
|
|
11
|
+
|
|
12
|
+
Conventions worth knowing before touching this:
|
|
13
|
+
* native sample rate 48 kHz; the codec is stereo internally, the public
|
|
14
|
+
interface is mono (decode averages the two channels).
|
|
15
|
+
* this export uses (batch, T, K) — TIME-major, codebook-LAST — for
|
|
16
|
+
``audio_codes``. ZeroTTS's AR loop produces (B, K, T), so decode transposes.
|
|
17
|
+
* every integer tensor here is **int32**, not int64. This is verified against
|
|
18
|
+
the graphs; do not "fix" it to int64 out of PyTorch habit.
|
|
19
|
+
|
|
20
|
+
Credit: MOSS-Audio-Tokenizer-Nano by the OpenMOSS team, Apache-2.0. See
|
|
21
|
+
https://github.com/OpenMOSS/MOSS-Audio-Tokenizer and the NOTICE file.
|
|
22
|
+
"""
|
|
23
|
+
|
|
24
|
+
from __future__ import annotations
|
|
25
|
+
|
|
26
|
+
import json
|
|
27
|
+
from pathlib import Path
|
|
28
|
+
|
|
29
|
+
import numpy as np
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class MossCodecDecoder:
|
|
33
|
+
"""Args:
|
|
34
|
+
codec_dir: directory holding the decoder graphs and
|
|
35
|
+
``codec_browser_onnx_meta.json``.
|
|
36
|
+
providers: onnxruntime execution providers.
|
|
37
|
+
intra_op_num_threads: per-session thread count.
|
|
38
|
+
"""
|
|
39
|
+
|
|
40
|
+
def __init__(
|
|
41
|
+
self,
|
|
42
|
+
codec_dir: str | Path,
|
|
43
|
+
providers: list[str] | None = None,
|
|
44
|
+
intra_op_num_threads: int = 4,
|
|
45
|
+
):
|
|
46
|
+
import onnxruntime as ort
|
|
47
|
+
|
|
48
|
+
codec_dir = Path(codec_dir)
|
|
49
|
+
meta_path = codec_dir / "codec_browser_onnx_meta.json"
|
|
50
|
+
if not meta_path.exists():
|
|
51
|
+
raise FileNotFoundError(
|
|
52
|
+
f"{codec_dir} has no codec_browser_onnx_meta.json — the vendored "
|
|
53
|
+
f"codec is missing from this model directory.")
|
|
54
|
+
meta = json.loads(meta_path.read_text())
|
|
55
|
+
self._meta = meta
|
|
56
|
+
|
|
57
|
+
cfg = meta["codec_config"]
|
|
58
|
+
self.sample_rate = int(cfg["sample_rate"])
|
|
59
|
+
self.num_channels = int(cfg["channels"])
|
|
60
|
+
self.frame_size = int(cfg["downsample_rate"])
|
|
61
|
+
self.frame_rate = self.sample_rate / self.frame_size
|
|
62
|
+
self.num_codebooks = int(cfg["num_quantizers"])
|
|
63
|
+
|
|
64
|
+
sess_options = ort.SessionOptions()
|
|
65
|
+
sess_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
|
|
66
|
+
sess_options.intra_op_num_threads = intra_op_num_threads
|
|
67
|
+
sess_options.inter_op_num_threads = 1
|
|
68
|
+
resolved = providers or ["CPUExecutionProvider"]
|
|
69
|
+
|
|
70
|
+
def _session(key: str):
|
|
71
|
+
return ort.InferenceSession(
|
|
72
|
+
str(codec_dir / meta["files"][key]),
|
|
73
|
+
sess_options=sess_options,
|
|
74
|
+
providers=resolved,
|
|
75
|
+
)
|
|
76
|
+
|
|
77
|
+
self._decode_full_sess = _session("decode_full")
|
|
78
|
+
self._decode_step_sess = _session("decode_step")
|
|
79
|
+
|
|
80
|
+
def decode(self, codes_bkt: np.ndarray) -> np.ndarray:
|
|
81
|
+
"""(B, K, T) int codes -> (B, T_audio) float32 mono at sample_rate."""
|
|
82
|
+
codes = np.asarray(codes_bkt)
|
|
83
|
+
if codes.ndim == 2: # (K, T) -> (1, K, T)
|
|
84
|
+
codes = codes[None, :, :]
|
|
85
|
+
codes_btk = codes.transpose(0, 2, 1).astype(np.int32) # -> (B, T, K)
|
|
86
|
+
lengths = np.array([codes_btk.shape[1]], dtype=np.int32)
|
|
87
|
+
audio, audio_lengths = self._decode_full_sess.run(
|
|
88
|
+
None, {"audio_codes": codes_btk, "audio_code_lengths": lengths}
|
|
89
|
+
)
|
|
90
|
+
n = int(audio_lengths.reshape(-1)[0])
|
|
91
|
+
return audio[:, :, :n].mean(axis=1).astype(np.float32)
|
|
92
|
+
|
|
93
|
+
def streaming_decoder(self) -> MossStreamingDecoder:
|
|
94
|
+
"""Open a stateful streaming decoder (keeps the causal decoder KV cache
|
|
95
|
+
across chunks). Call decode_chunk per chunk, then close."""
|
|
96
|
+
return MossStreamingDecoder(self)
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
class MossStreamingDecoder:
|
|
100
|
+
"""KV-cached streaming decode over decode_step.onnx, driven by the state
|
|
101
|
+
layout ``codec_browser_onnx_meta.json``'s "streaming_decode" section
|
|
102
|
+
describes: per-decoder transformer offsets plus per-layer attention caches
|
|
103
|
+
(key/value/position ring buffers). Use via
|
|
104
|
+
``MossCodecDecoder.streaming_decoder()``."""
|
|
105
|
+
|
|
106
|
+
def __init__(self, codec: MossCodecDecoder):
|
|
107
|
+
self._codec = codec
|
|
108
|
+
self._session = codec._decode_step_sess
|
|
109
|
+
streaming = codec._meta.get("streaming_decode", {})
|
|
110
|
+
self._transformer_specs = list(streaming.get("transformer_offsets", []))
|
|
111
|
+
self._attention_specs = list(streaming.get("attention_caches", []))
|
|
112
|
+
self._output_names = [o.name for o in self._session.get_outputs()]
|
|
113
|
+
self._state: dict = {}
|
|
114
|
+
self._reset_state()
|
|
115
|
+
|
|
116
|
+
def _reset_state(self) -> None:
|
|
117
|
+
self._state = {}
|
|
118
|
+
for spec in self._transformer_specs:
|
|
119
|
+
self._state[str(spec["input_name"])] = np.zeros(tuple(spec["shape"]), dtype=np.int32)
|
|
120
|
+
for spec in self._attention_specs:
|
|
121
|
+
self._state[str(spec["offset_input_name"])] = np.zeros(
|
|
122
|
+
tuple(spec["offset_shape"]), dtype=np.int32)
|
|
123
|
+
self._state[str(spec["cached_keys_input_name"])] = np.zeros(
|
|
124
|
+
tuple(spec["cache_shape"]), dtype=np.float32)
|
|
125
|
+
self._state[str(spec["cached_values_input_name"])] = np.zeros(
|
|
126
|
+
tuple(spec["cache_shape"]), dtype=np.float32)
|
|
127
|
+
# -1, not 0: position 0 is a real position, so a zero-filled ring
|
|
128
|
+
# buffer would read as "every slot holds frame 0".
|
|
129
|
+
self._state[str(spec["cached_positions_input_name"])] = np.full(
|
|
130
|
+
tuple(spec["positions_shape"]), -1, dtype=np.int32)
|
|
131
|
+
|
|
132
|
+
def decode_chunk(self, codes_bkt: np.ndarray) -> np.ndarray:
|
|
133
|
+
"""(1, K, n) int codes -> (1, chunk_samples) float32 mono."""
|
|
134
|
+
codes = np.asarray(codes_bkt)
|
|
135
|
+
if codes.ndim == 2:
|
|
136
|
+
codes = codes[None, :, :]
|
|
137
|
+
codes_btk = codes.transpose(0, 2, 1).astype(np.int32) # -> (1, n, K)
|
|
138
|
+
feeds = {
|
|
139
|
+
"audio_codes": codes_btk,
|
|
140
|
+
"audio_code_lengths": np.array([codes_btk.shape[1]], dtype=np.int32),
|
|
141
|
+
**self._state,
|
|
142
|
+
}
|
|
143
|
+
outputs = self._session.run(None, feeds)
|
|
144
|
+
named = dict(zip(self._output_names, outputs))
|
|
145
|
+
for spec in self._transformer_specs:
|
|
146
|
+
self._state[str(spec["input_name"])] = named[str(spec["output_name"])]
|
|
147
|
+
for spec in self._attention_specs:
|
|
148
|
+
for key in ("offset", "cached_keys", "cached_values", "cached_positions"):
|
|
149
|
+
self._state[str(spec[f"{key}_input_name"])] = named[str(spec[f"{key}_output_name"])]
|
|
150
|
+
n = int(named["audio_lengths"].reshape(-1)[0])
|
|
151
|
+
return named["audio"][:, :, :n].mean(axis=1).astype(np.float32)
|
|
152
|
+
|
|
153
|
+
def close(self) -> None:
|
|
154
|
+
self._reset_state()
|