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 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()