Luigi commited on
Commit
9994aaf
·
verified ·
1 Parent(s): f34bd6d

Add PrimeTTS v2-Stream-Clean: token-level input + v2-clean audio (RIGHT=16)

Browse files
v2streamclean_streaming/README.md ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # PrimeTTS v2-Stream-Clean — token-level streaming, v2-clean audio
2
+
3
+ The definitive streaming variant: **token-level band-attention encoder** (input
4
+ streams incrementally) **+ v2's non-causal clean vocoder** (no parasite noise).
5
+ Reconciles token-level input streaming with clean audio — the causal v2-Stream
6
+ sacrificed quality unnecessarily; this doesn't.
7
+
8
+ - `v2streamclean_enc.onnx` — text (x,tone,lang,x_lengths,noise_scale,length_scale) → z[1,192,T]. Band encoder, token-level. Run **once** per phrase.
9
+ - `v2streamclean_dec.onnx` — z[1,192,Tc] → wav[1,1,Tc·256]. Clean non-causal vocoder. Run **per chunk**, overlap-save.
10
+ - `onnx_stream.py` — reference runner (uses the right params).
11
+
12
+ **Streaming params: chunk = 24, left = 64, RIGHT = 16** (the clean non-causal
13
+ vocoder needs 16 future frames for bit-exact chunking; the causal one used 4).
14
+ 16 kHz, zh-TW + English.
15
+
16
+ ```python
17
+ from onnx_stream import StreamingTTS # RIGHT=16 baked in
18
+ tts = StreamingTTS("v2streamclean_enc.onnx", "v2streamclean_dec.onnx")
19
+ z = tts.encode(phone_ids, tone_ids, lang_ids) # once
20
+ for pcm in tts.stream(z): play(pcm) # per 24-frame chunk, clean audio
21
+ ```
22
+ Frontend (text→ids): g2pw bopomofo + g2p_en, 88 syms/6 tones/2 langs, add_blank.
23
+ sherpa-onnx: `OfflineTtsMbistftStreamModel(enc, dec, num_threads=2, right_lookahead=16)`.
24
+
25
+ License: Apache-2.0 · part of [Luigi/PrimeTTS](https://huggingface.co/Luigi/PrimeTTS).
v2streamclean_streaming/onnx_stream.py ADDED
@@ -0,0 +1,92 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Reference streaming TTS runner over the split v2-Stream ONNX models — the exact
2
+ orchestration a sherpa-onnx C++ runner should mirror.
3
+
4
+ enc.onnx : (x,tone,lang,x_lengths,noise_scale,length_scale) -> z[1,192,T] (once)
5
+ dec.onnx : z[1,192,Tc] -> wav[1,1,Tc*256] (per chunk)
6
+
7
+ Streaming = run enc once, then decode z in CHUNK-frame steps via OVERLAP-SAVE:
8
+ for chunk frames [a,b) decode z[:, :, a-LEFT : b+RIGHT] and keep the middle
9
+ (b-a)*256 samples. Bit-exact vs the monolithic model (validated: cos 1.000000,
10
+ maxerr ~1e-6). First audio arrives after enc + one chunk instead of the whole
11
+ utterance.
12
+
13
+ Usage:
14
+ python -m streaming.onnx_stream --enc <enc.onnx> --dec <dec.onnx> --ids <parity_inputs.json> [--i 0]
15
+ (or --text "..." with the g2pw frontend available)
16
+ """
17
+ from __future__ import annotations
18
+ import argparse, json, time
19
+ import numpy as np
20
+ import onnxruntime as ort
21
+
22
+ C, HOP, CHUNK, LEFT, RIGHT = 192, 256, 24, 64, 16 # non-causal clean vocoder needs 16 (causal was 4)
23
+
24
+
25
+ def _blank(seq):
26
+ o = [0] * (2 * len(seq) + 1)
27
+ o[1::2] = seq
28
+ return np.array([o], np.int64)
29
+
30
+
31
+ class StreamingTTS:
32
+ def __init__(self, enc_path, dec_path, threads=2):
33
+ so = ort.SessionOptions(); so.intra_op_num_threads = threads; so.inter_op_num_threads = 1
34
+ self.enc = ort.InferenceSession(enc_path, so, providers=["CPUExecutionProvider"])
35
+ self.dec = ort.InferenceSession(dec_path, so, providers=["CPUExecutionProvider"])
36
+
37
+ def encode(self, phone_ids, tone_ids, lang_ids, noise_scale=0.667, length_scale=1.0):
38
+ x, tn, lg = _blank(phone_ids), _blank(tone_ids), _blank(lang_ids)
39
+ return self.enc.run(None, {
40
+ "x": x, "tone": tn, "lang": lg,
41
+ "x_lengths": np.array([x.shape[1]], np.int64),
42
+ "noise_scale": np.array([noise_scale], np.float32),
43
+ "length_scale": np.array([length_scale], np.float32)})[0] # [1,192,T]
44
+
45
+ def stream(self, z):
46
+ """Yield audio chunks (np.float32) as z is decoded chunk-by-chunk."""
47
+ T = z.shape[2]
48
+ for a in range(0, T, CHUNK):
49
+ b = min(a + CHUNK, T); s0 = max(0, a - LEFT); e = min(T, b + RIGHT)
50
+ w = self.dec.run(None, {"z": z[:, :, s0:e]})[0].reshape(-1)
51
+ off = (a - s0) * HOP; keep = (b - a) * HOP
52
+ yield w[off:off + keep]
53
+
54
+ def synth(self, phone_ids, tone_ids, lang_ids, **kw):
55
+ z = self.encode(phone_ids, tone_ids, lang_ids, **kw)
56
+ return np.concatenate(list(self.stream(z))) if z.shape[2] else np.zeros(0, np.float32)
57
+
58
+
59
+ def main():
60
+ ap = argparse.ArgumentParser()
61
+ ap.add_argument("--enc", default="/home/luigi/mbvits_run/v2stream_split/v2stream_enc.onnx")
62
+ ap.add_argument("--dec", default="/home/luigi/mbvits_run/v2stream_split/v2stream_dec.onnx")
63
+ ap.add_argument("--ids", default="/home/luigi/mbvits_run/parity_inputs.json")
64
+ ap.add_argument("--i", type=int, default=0)
65
+ ap.add_argument("--text", default=None)
66
+ ap.add_argument("--out", default="/home/luigi/mbvits_run/onnx_stream_demo.wav")
67
+ ap.add_argument("--threads", type=int, default=2)
68
+ a = ap.parse_args()
69
+ tts = StreamingTTS(a.enc, a.dec, a.threads)
70
+
71
+ if a.text:
72
+ import sys; sys.path.insert(0, "/home/luigi/primetts-space")
73
+ import frontend_bopomofo as F
74
+ o = F.text_to_ids(a.text); p, t, l = o["phone_ids"], o["tone_ids"], o["lang_ids"]
75
+ else:
76
+ r = json.load(open(a.ids))["rows"][a.i]; p, t, l = r["phone_ids"], r["tone_ids"], r["lang_ids"]
77
+
78
+ t0 = time.perf_counter(); z = tts.encode(p, t, l); t_enc = time.perf_counter() - t0
79
+ chunks = []; tfirst = None
80
+ for c in tts.stream(z):
81
+ chunks.append(c)
82
+ if tfirst is None: tfirst = time.perf_counter() - t0
83
+ total = time.perf_counter() - t0
84
+ wav = np.concatenate(chunks); audio_s = len(wav) / 16000
85
+ pk = np.max(np.abs(wav)); wav = wav * (0.97 / pk) if pk > 1e-6 else wav
86
+ import soundfile as sf; sf.write(a.out, wav.astype(np.float32), 16000)
87
+ print(f"frames={z.shape[2]} audio={audio_s:.2f}s enc={t_enc*1e3:.0f}ms "
88
+ f"first-audio={tfirst*1e3:.0f}ms total={total*1e3:.0f}ms RTF={total/audio_s:.3f} -> {a.out}")
89
+
90
+
91
+ if __name__ == "__main__":
92
+ main()
v2streamclean_streaming/v2streamclean_dec.onnx ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:bfed5c287d651d46d48b63cec76b13b27732b2344069af77f02c458abf4801fc
3
+ size 54886617
v2streamclean_streaming/v2streamclean_enc.onnx ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:ddf7a0e2e2042899992ff669401b3ee2dce2474c763a47243ae78119f9effe20
3
+ size 55538543