qnguyen3 commited on
Commit
409d4fb
·
verified ·
1 Parent(s): d63631e

LGTM: PyTorch + ONNX weights and inference code

Browse files
README.md ADDED
@@ -0,0 +1,98 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ language: [en, es, pt, fr, de, it, sv, vi, ja, ko, id]
3
+ pipeline_tag: text-to-speech
4
+ tags: [text-to-speech, tts, voice-cloning, onnx, multilingual]
5
+ ---
6
+
7
+ # LGTM
8
+
9
+ Multilingual text-to-speech at 44.1 kHz with 10 built-in voices and zero-shot voice cloning.
10
+ Available in **PyTorch** and **ONNX** (ONNX Runtime needs no PyTorch).
11
+
12
+ **Languages:** English `en`, Spanish `es`, Portuguese `pt`, French `fr`, German `de`, Italian `it`,
13
+ Swedish `sv`, Vietnamese `vi`, Japanese `ja`, Korean `ko`, Indonesian `id`
14
+
15
+ **Built-in voices:** `F1` `F2` `F3` `F4` `F5` (female), `M1` `M2` `M3` `M4` `M5` (male)
16
+
17
+ ## Setup
18
+
19
+ ```bash
20
+ git clone https://huggingface.co/qnguyen3/LGTM
21
+ cd LGTM
22
+ pip install -r requirements.txt # PyTorch backend
23
+ pip install -r requirements-onnx.txt # ONNX backend
24
+ ```
25
+
26
+ ## PyTorch
27
+
28
+ ```python
29
+ from lgtm import LGTMTTS
30
+
31
+ tts = LGTMTTS.from_pretrained(".") # or "qnguyen3/LGTM" to download
32
+ wav = tts.synthesize("Xin chào, hôm nay trời đẹp quá!", lang="vi", voice="F1")
33
+ tts.save_wav(wav, "out.wav") # 44.1 kHz mono
34
+ ```
35
+
36
+ ## ONNX Runtime
37
+
38
+ ```python
39
+ from lgtm import LGTMOnnx
40
+
41
+ tts = LGTMOnnx.from_pretrained(".", use_gpu=False) # use_gpu=True with onnxruntime-gpu
42
+ wav = tts.synthesize("Bonjour à tous, comment allez-vous ?", lang="fr", voice="M1")
43
+ tts.save_wav(wav, "out.wav")
44
+ ```
45
+
46
+ ## Voice cloning
47
+
48
+ Give 5-15 seconds of clean speech; the voice can then speak any supported language.
49
+
50
+ ```python
51
+ voice = tts.clone_voice("reference.wav") # works with both backends
52
+ wav = tts.synthesize("This is my cloned voice.", lang="en", voice=voice)
53
+
54
+ from lgtm import save_voice_style # ONNX: from lgtm.onnx_inference import save_voice_style
55
+ save_voice_style("my_voice.json", voice) # reuse later: voice="my_voice.json"
56
+ ```
57
+
58
+ ## Command line
59
+
60
+ ```bash
61
+ python -m lgtm.cli --text "Hej! Hur mår du idag?" --lang sv --voice F2 --out out.wav
62
+ python -m lgtm.cli --text "안녕하세요" --lang ko --ref reference.wav --save_voice my_voice.json --out out.wav
63
+ python -m lgtm.cli --text "Selamat pagi" --lang id --voice M3 --backend onnx --out out.wav
64
+ ```
65
+
66
+ ## Options
67
+
68
+ | argument | default | |
69
+ |---|---|---|
70
+ | `lang` | `"en"` | language code (see above) |
71
+ | `voice` | `"F1"` | preset name, path to a voice `.json`, or `clone_voice()` output |
72
+ | `steps` | `8` | denoising steps (fewer = faster, e.g. 5) |
73
+ | `speed` | `1.05` | speaking rate (higher = faster) |
74
+ | `silence` | `0.3` | seconds of silence between sentences (long text is split automatically) |
75
+
76
+ ## Files
77
+
78
+ | path | contents |
79
+ |---|---|
80
+ | `pytorch/model.safetensors` | all weights (synthesis + voice cloning) |
81
+ | `onnx/text_encoder.onnx`, `duration_predictor.onnx`, `vector_estimator.onnx`, `vocoder.onnx` | synthesis graphs |
82
+ | `onnx/voice_encoder.onnx` | reference audio → voice style (cloning) |
83
+ | `voice_styles/*.json` | built-in voices |
84
+ | `config.json`, `unicode_indexer.json` | model config, text vocabulary |
85
+ | `lgtm/` | inference code (`inference.py` PyTorch, `onnx_inference.py` ONNX, `cli.py`) |
86
+
87
+ ### ONNX graph I/O (for custom runtimes)
88
+
89
+ | graph | inputs | outputs |
90
+ |---|---|---|
91
+ | `text_encoder` | `text_ids` int64 [B,T], `style_ttl` [B,50,256], `text_mask` [B,1,T] | `text_emb` [B,256,T] |
92
+ | `duration_predictor` | `text_ids`, `style_dp` [B,8,16], `text_mask` | `duration` [B] (seconds) |
93
+ | `vector_estimator` | `noisy_latent` [B,144,L], `text_emb`, `style_ttl`, `latent_mask` [B,1,L], `text_mask`, `current_step` [B], `total_step` [B] | `denoised_latent` [B,144,L] |
94
+ | `vocoder` | `latent` [B,144,L] | `wav` [B, 3072·L] |
95
+ | `voice_encoder` | `wav` [1,N] (44.1 kHz) | `style_ttl` [1,50,256], `style_dp` [1,8,16] |
96
+
97
+ Sampling loop: `L = ceil(duration·44100 / 3072)`, start from Gaussian noise masked by `latent_mask`,
98
+ call `vector_estimator` for `current_step = 0 … total_step-1`, then `vocoder`. See `lgtm/onnx_inference.py`.
config.json ADDED
@@ -0,0 +1,311 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "tts_version": "v1.7.3",
3
+ "split": "opensource-multilingual",
4
+ "ttl": {
5
+ "latent_dim": 24,
6
+ "chunk_compress_factor": 6,
7
+ "batch_expander": {
8
+ "n_batch_expand": 6
9
+ },
10
+ "normalizer": {
11
+ "scale": 0.25
12
+ },
13
+ "text_encoder": {
14
+ "n_langs": 0,
15
+ "lang_emb_dim": 0,
16
+ "text_embedder": {
17
+ "char_emb_dim": 256
18
+ },
19
+ "convnext": {
20
+ "idim": 256,
21
+ "ksz": 5,
22
+ "intermediate_dim": 1024,
23
+ "num_layers": 6,
24
+ "dilation_lst": [
25
+ 1,
26
+ 1,
27
+ 2,
28
+ 2,
29
+ 4,
30
+ 4
31
+ ]
32
+ },
33
+ "attn_encoder": {
34
+ "hidden_channels": 256,
35
+ "filter_channels": 1024,
36
+ "n_heads": 4,
37
+ "n_layers": 4,
38
+ "p_dropout": 0.0
39
+ },
40
+ "proj_out": {
41
+ "idim": 256,
42
+ "odim": 256
43
+ }
44
+ },
45
+ "flow_matching": {
46
+ "sig_min": 1e-08
47
+ },
48
+ "style_encoder": {
49
+ "proj_in": {
50
+ "ldim": 24,
51
+ "chunk_compress_factor": 6,
52
+ "odim": 256
53
+ },
54
+ "convnext": {
55
+ "idim": 256,
56
+ "ksz": 5,
57
+ "intermediate_dim": 1024,
58
+ "num_layers": 6,
59
+ "dilation_lst": [
60
+ 1,
61
+ 1,
62
+ 1,
63
+ 1,
64
+ 1,
65
+ 1
66
+ ]
67
+ },
68
+ "style_token_layer": {
69
+ "input_dim": 256,
70
+ "n_style": 50,
71
+ "style_key_dim": 256,
72
+ "style_value_dim": 256,
73
+ "prototype_dim": 256,
74
+ "n_units": 256,
75
+ "n_heads": 2
76
+ }
77
+ },
78
+ "speech_prompted_text_encoder": {
79
+ "text_dim": 256,
80
+ "style_dim": 256,
81
+ "n_units": 256,
82
+ "n_heads": 2
83
+ },
84
+ "uncond_masker": {
85
+ "prob_both_uncond": 0.04,
86
+ "prob_text_uncond": 0.01,
87
+ "std": 0.1,
88
+ "text_dim": 256,
89
+ "n_style": 50,
90
+ "style_key_dim": 256,
91
+ "style_value_dim": 256
92
+ },
93
+ "vector_field": {
94
+ "n_langs": 0,
95
+ "lang_emb_dim": 0,
96
+ "proj_in": {
97
+ "ldim": 24,
98
+ "chunk_compress_factor": 6,
99
+ "odim": 512
100
+ },
101
+ "time_encoder": {
102
+ "time_dim": 64,
103
+ "hdim": 256
104
+ },
105
+ "main_blocks": {
106
+ "n_blocks": 4,
107
+ "time_cond_layer": {
108
+ "idim": 512,
109
+ "time_dim": 64
110
+ },
111
+ "style_cond_layer": {
112
+ "idim": 512,
113
+ "style_dim": 256
114
+ },
115
+ "text_cond_layer": {
116
+ "idim": 512,
117
+ "text_dim": 256,
118
+ "n_heads": 8,
119
+ "n_units": 512,
120
+ "use_residual": true,
121
+ "rotary_base": 10000,
122
+ "rotary_scale": 10
123
+ },
124
+ "convnext_0": {
125
+ "idim": 512,
126
+ "ksz": 5,
127
+ "intermediate_dim": 2048,
128
+ "num_layers": 4,
129
+ "dilation_lst": [
130
+ 1,
131
+ 2,
132
+ 4,
133
+ 8
134
+ ]
135
+ },
136
+ "convnext_1": {
137
+ "idim": 512,
138
+ "ksz": 5,
139
+ "intermediate_dim": 2048,
140
+ "num_layers": 1,
141
+ "dilation_lst": [
142
+ 1
143
+ ]
144
+ },
145
+ "convnext_2": {
146
+ "idim": 512,
147
+ "ksz": 5,
148
+ "intermediate_dim": 2048,
149
+ "num_layers": 1,
150
+ "dilation_lst": [
151
+ 1
152
+ ]
153
+ }
154
+ },
155
+ "last_convnext": {
156
+ "idim": 512,
157
+ "ksz": 5,
158
+ "intermediate_dim": 2048,
159
+ "num_layers": 4,
160
+ "dilation_lst": [
161
+ 1,
162
+ 1,
163
+ 1,
164
+ 1
165
+ ]
166
+ },
167
+ "proj_out": {
168
+ "idim": 512,
169
+ "chunk_compress_factor": 6,
170
+ "ldim": 24
171
+ }
172
+ }
173
+ },
174
+ "ae": {
175
+ "sample_rate": 44100,
176
+ "n_delay": 0,
177
+ "base_chunk_size": 512,
178
+ "chunk_compress_factor": 1,
179
+ "ldim": 24,
180
+ "encoder": {
181
+ "spec_processor": {
182
+ "n_fft": 2048,
183
+ "win_length": 2048,
184
+ "hop_length": 512,
185
+ "n_mels": 228,
186
+ "sample_rate": 44100,
187
+ "eps": 1e-05,
188
+ "norm_mean": 0.0,
189
+ "norm_std": 1.0
190
+ },
191
+ "ksz_init": 7,
192
+ "ksz": 7,
193
+ "num_layers": 10,
194
+ "dilation_lst": [
195
+ 1,
196
+ 1,
197
+ 1,
198
+ 1,
199
+ 1,
200
+ 1,
201
+ 1,
202
+ 1,
203
+ 1,
204
+ 1
205
+ ],
206
+ "intermediate_dim": 2048,
207
+ "idim": 1253,
208
+ "hdim": 512,
209
+ "odim": 24
210
+ },
211
+ "decoder": {
212
+ "ksz_init": 7,
213
+ "ksz": 7,
214
+ "num_layers": 10,
215
+ "dilation_lst": [
216
+ 1,
217
+ 2,
218
+ 4,
219
+ 1,
220
+ 2,
221
+ 4,
222
+ 1,
223
+ 1,
224
+ 1,
225
+ 1
226
+ ],
227
+ "intermediate_dim": 2048,
228
+ "idim": 24,
229
+ "hdim": 512,
230
+ "head": {
231
+ "idim": 512,
232
+ "hdim": 2048,
233
+ "odim": 512,
234
+ "ksz": 3
235
+ }
236
+ }
237
+ },
238
+ "dp": {
239
+ "latent_dim": 24,
240
+ "chunk_compress_factor": 6,
241
+ "normalizer": {
242
+ "scale": 1.0
243
+ },
244
+ "sentence_encoder": {
245
+ "char_emb_dim": 64,
246
+ "text_embedder": {
247
+ "char_emb_dim": 64
248
+ },
249
+ "convnext": {
250
+ "idim": 64,
251
+ "ksz": 5,
252
+ "intermediate_dim": 256,
253
+ "num_layers": 6,
254
+ "dilation_lst": [
255
+ 1,
256
+ 1,
257
+ 1,
258
+ 1,
259
+ 1,
260
+ 1
261
+ ]
262
+ },
263
+ "attn_encoder": {
264
+ "hidden_channels": 64,
265
+ "filter_channels": 256,
266
+ "n_heads": 2,
267
+ "n_layers": 2,
268
+ "p_dropout": 0.0
269
+ },
270
+ "proj_out": {
271
+ "idim": 64,
272
+ "odim": 64
273
+ }
274
+ },
275
+ "style_encoder": {
276
+ "proj_in": {
277
+ "ldim": 24,
278
+ "chunk_compress_factor": 6,
279
+ "odim": 64
280
+ },
281
+ "convnext": {
282
+ "idim": 64,
283
+ "ksz": 5,
284
+ "intermediate_dim": 256,
285
+ "num_layers": 4,
286
+ "dilation_lst": [
287
+ 1,
288
+ 1,
289
+ 1,
290
+ 1
291
+ ]
292
+ },
293
+ "style_token_layer": {
294
+ "input_dim": 64,
295
+ "n_style": 8,
296
+ "style_key_dim": 0,
297
+ "style_value_dim": 16,
298
+ "prototype_dim": 64,
299
+ "n_units": 64,
300
+ "n_heads": 2
301
+ }
302
+ },
303
+ "predictor": {
304
+ "sentence_dim": 64,
305
+ "n_style": 8,
306
+ "style_dim": 16,
307
+ "hdim": 128,
308
+ "n_layer": 2
309
+ }
310
+ }
311
+ }
lgtm/__init__.py ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """LGTM text-to-speech. `LGTMTTS` = PyTorch backend, `LGTMOnnx` = ONNX Runtime backend (no torch needed)."""
2
+
3
+
4
+ def __getattr__(name): # lazy imports: the ONNX backend must not require torch
5
+ if name in ("LGTMTTS", "load_voice_style", "save_voice_style"):
6
+ from . import inference
7
+ return getattr(inference, name)
8
+ if name == "LGTM":
9
+ from .model import LGTM
10
+ return LGTM
11
+ if name == "LGTMOnnx":
12
+ from .onnx_inference import LGTMOnnx
13
+ return LGTMOnnx
14
+ raise AttributeError(name)
lgtm/cli.py ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Command line: python -m lgtm.cli --text "Xin chào" --lang vi --voice F1 --out out.wav [--ref ref.wav] [--backend onnx]"""
2
+ import argparse
3
+
4
+
5
+ def main():
6
+ ap = argparse.ArgumentParser(description="LGTM text-to-speech")
7
+ ap.add_argument("--text", required=True)
8
+ ap.add_argument("--lang", default="en", help="en es pt fr de it sv vi ja ko id")
9
+ ap.add_argument("--voice", default="F1", help="preset (F1-F5, M1-M5) or voice .json")
10
+ ap.add_argument("--ref", default=None, help="reference .wav to clone the voice from (overrides --voice)")
11
+ ap.add_argument("--save_voice", default=None, help="save the cloned voice to this .json")
12
+ ap.add_argument("--out", default="out.wav")
13
+ ap.add_argument("--model", default=".", help="local model dir or Hugging Face repo id")
14
+ ap.add_argument("--backend", choices=["torch", "onnx"], default="torch")
15
+ ap.add_argument("--steps", type=int, default=8)
16
+ ap.add_argument("--speed", type=float, default=1.05)
17
+ ap.add_argument("--gpu", action="store_true", help="(onnx) use CUDAExecutionProvider")
18
+ a = ap.parse_args()
19
+ if a.backend == "onnx":
20
+ from .onnx_inference import LGTMOnnx, save_voice_style
21
+ tts = LGTMOnnx.from_pretrained(a.model, use_gpu=a.gpu)
22
+ else:
23
+ from .inference import LGTMTTS, save_voice_style
24
+ tts = LGTMTTS.from_pretrained(a.model)
25
+ voice = tts.clone_voice(a.ref) if a.ref else a.voice
26
+ if a.ref and a.save_voice:
27
+ save_voice_style(a.save_voice, voice)
28
+ wav = tts.synthesize(a.text, lang=a.lang, voice=voice, steps=a.steps, speed=a.speed)
29
+ tts.save_wav(wav, a.out)
30
+ print(f"wrote {a.out} ({len(wav) / 44100:.2f} s)")
31
+
32
+
33
+ if __name__ == "__main__":
34
+ main()
lgtm/export_onnx.py ADDED
@@ -0,0 +1,170 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Export LGTM to ONNX (5 graphs) and verify each against PyTorch.
2
+
3
+ python -m lgtm.export_onnx --model_dir . --out onnx
4
+
5
+ Graphs (all with dynamic batch / length):
6
+ text_encoder.onnx text_ids[B,T] int64, style_ttl[B,50,256], text_mask[B,1,T] -> text_emb[B,256,T]
7
+ duration_predictor.onnx text_ids[B,T], style_dp[B,8,16], text_mask[B,1,T] -> duration[B] (seconds)
8
+ vector_estimator.onnx noisy_latent[B,144,L], text_emb, style_ttl, latent_mask[B,1,L],
9
+ text_mask, current_step[B], total_step[B] -> denoised_latent[B,144,L]
10
+ vocoder.onnx latent[B,144,L] -> wav[B, 3072*L] (44.1 kHz)
11
+ voice_encoder.onnx wav[1,N] (44.1 kHz, prepared reference) -> style_ttl[1,50,256], style_dp[1,8,16]
12
+ """
13
+ import argparse
14
+ import json
15
+ import os
16
+
17
+ import numpy as np
18
+ import onnxruntime as ort
19
+ import torch
20
+ from safetensors.torch import load_file
21
+
22
+ from .model import LGTM
23
+ from .strip_onnx import strip
24
+
25
+
26
+ class TextEncoder(torch.nn.Module):
27
+ def __init__(self, m):
28
+ super().__init__()
29
+ self.m = m
30
+
31
+ def forward(self, text_ids, style_ttl, text_mask):
32
+ return self.m.ttl.encode_text(text_ids, text_mask, style_ttl)
33
+
34
+
35
+ class DurationPredictor(torch.nn.Module):
36
+ def __init__(self, m):
37
+ super().__init__()
38
+ self.m = m
39
+
40
+ def forward(self, text_ids, style_dp, text_mask):
41
+ return self.m.dp(text_ids, text_mask, style_dp)
42
+
43
+
44
+ class VectorEstimator(torch.nn.Module):
45
+ def __init__(self, m, cfg_scale=4.0):
46
+ super().__init__()
47
+ self.m, self.cfg_scale = m, cfg_scale
48
+
49
+ def forward(self, noisy_latent, text_emb, style_ttl, latent_mask, text_mask, current_step, total_step):
50
+ return self.m.ttl.euler_step(noisy_latent, current_step, total_step, text_emb, text_mask, style_ttl, latent_mask, self.cfg_scale)
51
+
52
+
53
+ class Vocoder(torch.nn.Module):
54
+ def __init__(self, m):
55
+ super().__init__()
56
+ self.m = m
57
+
58
+ def forward(self, latent):
59
+ return self.m.ae.decode_ttl(latent)
60
+
61
+
62
+ class VoiceEncoder(torch.nn.Module):
63
+ def __init__(self, m):
64
+ super().__init__()
65
+ self.m = m
66
+
67
+ def forward(self, wav):
68
+ lat = self.m.ae.encode_ttl(wav)
69
+ mask = torch.ones_like(lat[:, :1])
70
+ return self.m.ttl.style_encoder(lat, mask), self.m.dp.style_encoder(lat, mask)
71
+
72
+
73
+ def export(module, args, path, input_names, output_names, dynamic_shapes):
74
+ prog = torch.onnx.export(module, args, path, input_names=input_names, output_names=output_names,
75
+ dynamic_shapes=dynamic_shapes, dynamo=True, opset_version=20, optimize=True, external_data=False)
76
+ return prog
77
+
78
+
79
+ def check(path, module, feeds, names, atol_rel=1e-3):
80
+ sess = ort.InferenceSession(path, providers=["CPUExecutionProvider"])
81
+ outs = sess.run(None, {k: v.numpy() for k, v in feeds.items()})
82
+ with torch.no_grad():
83
+ ref = module(*feeds.values())
84
+ ref = ref if isinstance(ref, (tuple, list)) else (ref,)
85
+ worst = 0.0
86
+ for o, r, n in zip(outs, ref, names):
87
+ r = r.numpy()
88
+ rel = np.abs(o - r).max() / (np.abs(r).max() + 1e-9)
89
+ worst = max(worst, rel)
90
+ print(f" {n:16s} shape {tuple(o.shape)} max rel diff {rel:.2e}")
91
+ assert worst < atol_rel, f"{path}: mismatch {worst}"
92
+
93
+
94
+ def main():
95
+ ap = argparse.ArgumentParser()
96
+ ap.add_argument("--model_dir", default=".")
97
+ ap.add_argument("--out", default="onnx")
98
+ args = ap.parse_args()
99
+ os.makedirs(args.out, exist_ok=True)
100
+ cfg = json.load(open(os.path.join(args.model_dir, "config.json")))
101
+ m = LGTM(cfg)
102
+ m.load_state_dict(load_file(os.path.join(args.model_dir, "pytorch", "model.safetensors")))
103
+ m.eval()
104
+ torch.manual_seed(0)
105
+ B, T, L = 3, 57, 63 # B must not equal any model dim (e.g. 2 attention heads) or the exporter fixes it
106
+ D = torch.export.Dim
107
+ b, t, l, n = D("batch", min=1, max=64), D("text_len", min=2, max=1000), D("latent_len", min=2, max=1000), D("samples", min=15000, max=44100 * 60)
108
+
109
+ def masks(T_, L_, B_=B):
110
+ tm = torch.ones(B_, 1, T_)
111
+ tm[-1, :, T_ - 9:] = 0
112
+ lm = torch.ones(B_, 1, L_)
113
+ lm[-1, :, L_ - 11:] = 0
114
+ return tm, lm
115
+
116
+ ids = torch.randint(0, 8000, (B, T))
117
+ tm, lm = masks(T, L)
118
+ s_ttl = torch.nn.functional.normalize(torch.randn(B, 50, 256), dim=-1)
119
+ s_dp = torch.nn.functional.normalize(torch.randn(B, 8, 16), dim=-1)
120
+ with torch.no_grad():
121
+ te = m.ttl.encode_text(ids, tm, s_ttl)
122
+ x = torch.randn(B, 144, L) * lm
123
+
124
+ jobs = [
125
+ ("text_encoder", TextEncoder(m), {"text_ids": ids, "style_ttl": s_ttl, "text_mask": tm}, ["text_emb"],
126
+ {"text_ids": {0: b, 1: t}, "style_ttl": {0: b}, "text_mask": {0: b, 2: t}}),
127
+ ("duration_predictor", DurationPredictor(m), {"text_ids": ids, "style_dp": s_dp, "text_mask": tm}, ["duration"],
128
+ {"text_ids": {0: b, 1: t}, "style_dp": {0: b}, "text_mask": {0: b, 2: t}}),
129
+ ("vector_estimator", VectorEstimator(m),
130
+ {"noisy_latent": x, "text_emb": te, "style_ttl": s_ttl, "latent_mask": lm, "text_mask": tm,
131
+ "current_step": torch.full((B,), 3.0), "total_step": torch.full((B,), 8.0)}, ["denoised_latent"],
132
+ {"noisy_latent": {0: b, 2: l}, "text_emb": {0: b, 2: t}, "style_ttl": {0: b}, "latent_mask": {0: b, 2: l},
133
+ "text_mask": {0: b, 2: t}, "current_step": {0: b}, "total_step": {0: b}}),
134
+ ("vocoder", Vocoder(m), {"latent": x}, ["wav"], {"latent": {0: b, 2: l}}),
135
+ ("voice_encoder", VoiceEncoder(m), {"wav": torch.randn(1, 44100 * 4) * 0.05}, ["style_ttl", "style_dp"],
136
+ {"wav": {1: n}}),
137
+ ]
138
+ for name, mod, feeds, outs, dyn in jobs:
139
+ path = os.path.join(args.out, f"{name}.onnx")
140
+ print(f"exporting {name} ...", flush=True)
141
+ export(mod, tuple(feeds.values()), path, list(feeds), outs, dyn)
142
+ print(f" check (export shapes):")
143
+ check(path, mod, feeds, outs)
144
+ # different lengths than at export time, to prove the graph is really dynamic
145
+ B2 = 1 if name != "vocoder" else 5 # a different batch size than at export time
146
+ if name in ("text_encoder", "duration_predictor"):
147
+ T2 = 23
148
+ tm2, _ = masks(T2, L, B2)
149
+ second = list(feeds)[1]
150
+ f2 = {"text_ids": torch.randint(0, 8000, (B2, T2)), second: feeds[second][:B2], "text_mask": tm2}
151
+ elif name == "vector_estimator":
152
+ B2, T2, L2 = 5, 31, 140
153
+ tm2, lm2 = masks(T2, L2, B2)
154
+ st2 = torch.nn.functional.normalize(torch.randn(B2, 50, 256), dim=-1)
155
+ with torch.no_grad():
156
+ te2 = m.ttl.encode_text(torch.randint(0, 8000, (B2, T2)), tm2, st2)
157
+ f2 = {"noisy_latent": torch.randn(B2, 144, L2) * lm2, "text_emb": te2, "style_ttl": st2, "latent_mask": lm2,
158
+ "text_mask": tm2, "current_step": torch.zeros(B2), "total_step": torch.full((B2,), 8.0)}
159
+ elif name == "vocoder":
160
+ f2 = {"latent": torch.randn(B2, 144, 17) * 0.27}
161
+ else:
162
+ f2 = {"wav": torch.randn(1, 44100 * 9) * 0.05}
163
+ print(f" check (new shapes):")
164
+ check(path, mod, f2, outs)
165
+ strip(path) # drop exporter debug metadata (stack traces with local paths)
166
+ print(f" {os.path.getsize(path) / 1e6:.1f} MB", flush=True)
167
+
168
+
169
+ if __name__ == "__main__":
170
+ main()
lgtm/inference.py ADDED
@@ -0,0 +1,151 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """PyTorch inference for LGTM.
2
+
3
+ from lgtm import LGTMTTS
4
+ tts = LGTMTTS.from_pretrained("path/or/hf-repo") # downloads from the Hub if needed
5
+ wav = tts.synthesize("Xin chào!", lang="vi", voice="F1")
6
+ tts.save_wav(wav, "out.wav")
7
+ voice = tts.clone_voice("reference.wav") # zero-shot voice from 5-15 s of audio
8
+ wav = tts.synthesize("Hello there.", lang="en", voice=voice)
9
+ """
10
+ import json
11
+ import os
12
+ import re
13
+
14
+ import numpy as np
15
+ import soundfile as sf
16
+ import torch
17
+ import torch.nn.functional as F
18
+ import torchaudio
19
+ from safetensors.torch import load_file
20
+
21
+ from .model import LGTM
22
+ from .text import AVAILABLE_LANGS, TextProcessor
23
+
24
+ SAMPLE_RATE = 44100
25
+
26
+
27
+ def load_voice_style(path, device="cpu"):
28
+ d = json.load(open(path))
29
+ ttl = torch.tensor(np.array(d["style_ttl"]["data"], np.float32).reshape(d["style_ttl"]["dims"]))
30
+ dp = torch.tensor(np.array(d["style_dp"]["data"], np.float32).reshape(d["style_dp"]["dims"]))
31
+ return ttl.to(device), dp.to(device)
32
+
33
+
34
+ def save_voice_style(path, voice):
35
+ ttl, dp = voice
36
+ d = {"style_ttl": {"data": ttl.detach().float().cpu().numpy().tolist(), "dims": list(ttl.shape), "type": "float32"},
37
+ "style_dp": {"data": dp.detach().float().cpu().numpy().tolist(), "dims": list(dp.shape), "type": "float32"}}
38
+ json.dump(d, open(path, "w"))
39
+
40
+
41
+ def _prepare_reference(w, top_db=40.0, pad_s=0.1, target_rms_db=-23.0, peak=0.95):
42
+ """Trim leading/trailing silence and normalise loudness (same as the training data)."""
43
+ hop = SAMPLE_RATE // 100
44
+ n = len(w) // hop
45
+ if n >= 3:
46
+ db = 10 * torch.log10(w[: n * hop].view(n, hop).pow(2).mean(1) + 1e-10)
47
+ active = torch.nonzero(db > db.max() - top_db).flatten()
48
+ if len(active):
49
+ pad = int(pad_s * SAMPLE_RATE)
50
+ w = w[max(int(active[0]) * hop - pad, 0): min((int(active[-1]) + 1) * hop + pad, len(w))]
51
+ w = w * (10 ** (target_rms_db / 20) / (w.pow(2).mean().sqrt() + 1e-6))
52
+ m = w.abs().max()
53
+ return w * (peak / m) if m > peak else w
54
+
55
+
56
+ def split_text(text, max_len=300):
57
+ """Split long text into sentence chunks of at most max_len characters."""
58
+ chunks = []
59
+ for para in [p.strip() for p in re.split(r"\n\s*\n+", text.strip()) if p.strip()]:
60
+ cur = ""
61
+ for sent in re.split(r"(?<=[.!?。!?])\s+", para):
62
+ if cur and len(cur) + len(sent) + 1 > max_len:
63
+ chunks.append(cur)
64
+ cur = sent
65
+ else:
66
+ cur = f"{cur} {sent}".strip()
67
+ if cur:
68
+ chunks.append(cur)
69
+ return chunks
70
+
71
+
72
+ class LGTMTTS:
73
+ def __init__(self, model_dir, device=None):
74
+ self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
75
+ cfg = json.load(open(os.path.join(model_dir, "config.json")))
76
+ self.model = LGTM(cfg)
77
+ self.model.load_state_dict(load_file(os.path.join(model_dir, "pytorch", "model.safetensors")))
78
+ self.model.to(self.device).eval()
79
+ self.tp = TextProcessor(os.path.join(model_dir, "unicode_indexer.json"))
80
+ self.voice_dir = os.path.join(model_dir, "voice_styles")
81
+
82
+ @classmethod
83
+ def from_pretrained(cls, repo_or_dir, device=None):
84
+ if not os.path.isdir(repo_or_dir):
85
+ from huggingface_hub import snapshot_download
86
+
87
+ repo_or_dir = snapshot_download(repo_or_dir, allow_patterns=["config.json", "unicode_indexer.json", "voice_styles/*", "pytorch/*"])
88
+ return cls(repo_or_dir, device)
89
+
90
+ @property
91
+ def voices(self):
92
+ return sorted(f[:-5] for f in os.listdir(self.voice_dir) if f.endswith(".json"))
93
+
94
+ def _voice(self, voice):
95
+ if isinstance(voice, str):
96
+ path = voice if voice.endswith(".json") else os.path.join(self.voice_dir, f"{voice}.json")
97
+ return load_voice_style(path, self.device)
98
+ return voice[0].to(self.device), voice[1].to(self.device)
99
+
100
+ @torch.no_grad()
101
+ def clone_voice(self, wav_path, max_seconds=15.0):
102
+ """Voice style from a reference recording (5-15 s of clean speech works best)."""
103
+ wav, sr = sf.read(wav_path, dtype="float32")
104
+ if wav.ndim == 2:
105
+ wav = wav.mean(1)
106
+ w = torch.from_numpy(wav)
107
+ if sr != SAMPLE_RATE:
108
+ w = torchaudio.functional.resample(w, sr, SAMPLE_RATE)
109
+ w = _prepare_reference(w)[: int(max_seconds * SAMPLE_RATE)]
110
+ lat = self.model.ae.encode_ttl(w[None].to(self.device))
111
+ mask = torch.ones(1, 1, lat.shape[-1], device=self.device)
112
+ return self.model.ttl.style_encoder(lat, mask), self.model.dp.style_encoder(lat, mask)
113
+
114
+ @torch.no_grad()
115
+ def _synth_batch(self, texts, langs, voice, steps, speed, cfg_scale):
116
+ s_ttl, s_dp = voice
117
+ b = len(texts)
118
+ ids, mask = self.tp.batch(texts, langs, device=self.device)
119
+ s_ttl, s_dp = s_ttl.expand(b, -1, -1), s_dp.expand(b, -1, -1)
120
+ m = self.model
121
+ dur = m.dp(ids, mask, s_dp) / speed
122
+ text_emb = m.ttl.encode_text(ids, mask, s_ttl)
123
+ chunk = m.ae.hop * m.ae.ccf
124
+ wav_len = (dur * SAMPLE_RATE).long()
125
+ lat_len = (wav_len + chunk - 1) // chunk
126
+ L = int(lat_len.max())
127
+ lmask = (torch.arange(L, device=self.device)[None] < lat_len[:, None]).float().unsqueeze(1)
128
+ x = torch.randn(b, m.ae.ldim * m.ae.ccf, L, device=self.device) * lmask
129
+ total = torch.full((b,), float(steps), device=self.device)
130
+ for i in range(steps):
131
+ x = m.ttl.euler_step(x, torch.full((b,), float(i), device=self.device), total, text_emb, mask, s_ttl, lmask, cfg_scale)
132
+ wav = m.ae.decode_ttl(x)
133
+ return [wav[k, : wav_len[k]].float().cpu().numpy() for k in range(b)]
134
+
135
+ def synthesize(self, text, lang="en", voice="F1", steps=8, speed=1.05, cfg_scale=4.0, silence=0.3):
136
+ """Text -> 44.1 kHz float32 waveform. Long text is split into sentence chunks.
137
+ voice: preset name (F1..F5, M1..M5), a .json path, or the output of clone_voice()."""
138
+ if lang not in AVAILABLE_LANGS:
139
+ raise ValueError(f"unsupported language {lang!r}")
140
+ v = self._voice(voice)
141
+ chunks = split_text(text, 120 if lang in ("ja", "ko") else 300)
142
+ wavs = self._synth_batch(chunks, [lang] * len(chunks), v, steps, speed, cfg_scale)
143
+ gap = np.zeros(int(silence * SAMPLE_RATE), np.float32)
144
+ out = []
145
+ for i, w in enumerate(wavs):
146
+ out += [w] + ([gap] if i < len(wavs) - 1 else [])
147
+ return np.concatenate(out)
148
+
149
+ @staticmethod
150
+ def save_wav(wav, path):
151
+ sf.write(path, wav, SAMPLE_RATE)
lgtm/model.py ADDED
@@ -0,0 +1,441 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """LGTM text-to-speech model (44.1 kHz, 11 languages, zero-shot voice cloning).
2
+
3
+ LGTM
4
+ ttl text-to-latent: text encoder, style encoder, flow-matching vector field
5
+ ae speech autoencoder: encoder (audio -> latent) and decoder (latent -> 44.1 kHz audio)
6
+ dp utterance-level duration predictor (+ its style encoder)
7
+
8
+ Latents: raw AE latent (B, 24, T_f) at 44100/512 Hz is normalised and 6 frames are stacked into
9
+ channels -> (B, 144, T_f / 6); the flow-matching model works in that space.
10
+ """
11
+ import torch
12
+ import torch.nn as nn
13
+ import torch.nn.functional as F
14
+
15
+ from .modules import (
16
+ AttnEncoder,
17
+ ChannelLayerNorm,
18
+ CharEmbedder,
19
+ ConvNeXtStack,
20
+ Linear,
21
+ PaddedConv1d,
22
+ RotaryCrossAttention,
23
+ StyleAttention,
24
+ TimeEncoder,
25
+ )
26
+
27
+ N_VOCAB = 8322
28
+
29
+
30
+ # ==========================================================================
31
+ # Text-to-latent
32
+ # ==========================================================================
33
+ class TextEncoder(nn.Module):
34
+ """char emb -> ConvNeXt x6 -> relpos transformer x4, with a long residual."""
35
+
36
+ def __init__(self, cfg):
37
+ super().__init__()
38
+ self.text_embedder = CharEmbedder(N_VOCAB, cfg["text_embedder"]["char_emb_dim"])
39
+ self.convnext = ConvNeXtStack(**cfg["convnext"])
40
+ self.attn_encoder = AttnEncoder(**cfg["attn_encoder"])
41
+ # cfg["proj_out"] is an identity projection.
42
+
43
+ def forward(self, text_ids, text_mask):
44
+ x = self.text_embedder(text_ids, text_mask)
45
+ x = self.convnext(x, text_mask)
46
+ x = self.attn_encoder(x, text_mask) + x
47
+ return x * text_mask
48
+
49
+
50
+ class SpeechPromptedTextEncoder(nn.Module):
51
+ """Two cross-attentions from text to the 50 style tokens; both residuals
52
+ are taken w.r.t. the *original* text features (as in the graph)."""
53
+
54
+ def __init__(self, text_dim, style_dim, n_units, n_heads):
55
+ super().__init__()
56
+ self.attention1 = StyleAttention(text_dim, style_dim, style_dim, n_units, n_heads, text_dim)
57
+ self.attention2 = StyleAttention(text_dim, style_dim, style_dim, n_units, n_heads, text_dim)
58
+ self.norm = ChannelLayerNorm(text_dim)
59
+
60
+ def forward(self, x, text_mask, style_key, style_value):
61
+ xt = x.transpose(1, 2)
62
+ m = text_mask.transpose(1, 2)
63
+ h = self.attention1(xt, style_key, style_value, m) * m + xt
64
+ h = self.attention2(h, style_key, style_value, m) * m + xt
65
+ return self.norm.norm(h).transpose(1, 2) * text_mask
66
+
67
+
68
+ class StyleTokenLayer(nn.Module):
69
+ """Reference encoder pooling -> n_style tokens.
70
+
71
+
72
+ """
73
+
74
+ def __init__(self, input_dim, n_style, style_key_dim, style_value_dim, prototype_dim, n_units, n_heads):
75
+ super().__init__()
76
+ if style_key_dim > 0:
77
+ self.style_key = nn.Parameter(torch.randn(1, n_style, style_key_dim) * 0.02)
78
+ else:
79
+ self.style_key = None
80
+ self.prototype = nn.Parameter(torch.randn(1, n_style, prototype_dim) * 0.02)
81
+ self.attention = StyleAttention(prototype_dim, input_dim, input_dim, n_units, n_heads, style_value_dim)
82
+ # Output = normalize(base + delta): `base` is the mean voice style, the network predicts the
83
+ # voice-specific deviation.
84
+ self.base = nn.Parameter(torch.zeros(1, n_style, style_value_dim))
85
+ nn.init.zeros_(self.attention.out_fc.linear.weight)
86
+ nn.init.zeros_(self.attention.out_fc.linear.bias)
87
+
88
+ def forward(self, h, mask):
89
+ """h: (B, C, T) encoded reference, mask: (B, 1, T) -> (B, n_style, style_value_dim)."""
90
+ b = h.shape[0]
91
+ ht = h.transpose(1, 2)
92
+ q = self.prototype.expand(b, -1, -1)
93
+ # mask padded reference frames by pushing their keys to a neutral value
94
+ # and removing their values; StyleAttention has no key mask of its own.
95
+ return self._masked_attention(q, ht, mask.transpose(1, 2))
96
+
97
+ def _masked_attention(self, q, kv, kv_mask):
98
+ att = self.attention
99
+ qh = att._heads(att.W_query(q))
100
+ kh = torch.tanh(att._heads(att.W_key(kv)))
101
+ vh = att._heads(att.W_value(kv))
102
+ scores = torch.matmul(qh, kh.transpose(-1, -2)) / (att.n_units ** 0.5)
103
+ scores = scores.masked_fill(kv_mask.transpose(1, 2).unsqueeze(0) == 0, float("-inf"))
104
+ out = torch.matmul(F.softmax(scores, dim=-1), vh)
105
+ out = self.base + att.out_fc(torch.cat(out.unbind(0), dim=-1))
106
+ # voice styles (ttl 50x256 and dp 8x16) are per-token unit vectors
107
+ return F.normalize(out, dim=-1)
108
+
109
+
110
+ class StyleEncoder(nn.Module):
111
+ """Reference latent (B,144,T) -> style tokens."""
112
+
113
+ def __init__(self, cfg):
114
+ super().__init__()
115
+ p = cfg["proj_in"]
116
+ self.proj_in = PaddedConv1d(p["ldim"] * p["chunk_compress_factor"], p["odim"], 1)
117
+ self.convnext = ConvNeXtStack(**cfg["convnext"])
118
+ self.style_token_layer = StyleTokenLayer(**cfg["style_token_layer"])
119
+
120
+ def forward(self, latent, mask):
121
+ h = self.proj_in(latent) * mask
122
+ h = self.convnext(h, mask)
123
+ return self.style_token_layer(h, mask)
124
+
125
+
126
+ class UncondMasker(nn.Module):
127
+ """Learned null tokens for classifier-free guidance."""
128
+
129
+ def __init__(self, text_dim, n_style, style_key_dim, style_value_dim, **_):
130
+ super().__init__()
131
+ self.text_special_token = nn.Parameter(torch.zeros(1, text_dim, 1))
132
+ self.style_key_special_token = nn.Parameter(torch.zeros(1, n_style, style_key_dim))
133
+ self.style_value_special_token = nn.Parameter(torch.zeros(1, n_style, style_value_dim))
134
+
135
+
136
+ class TimeCondBlock(nn.Module):
137
+ def __init__(self, idim, time_dim):
138
+ super().__init__()
139
+ self.linear = Linear(time_dim, idim)
140
+
141
+ def forward(self, x, mask, t_emb):
142
+ return (x + self.linear(t_emb).unsqueeze(-1)) * mask
143
+
144
+
145
+ class TextCondBlock(nn.Module):
146
+ def __init__(self, idim, text_dim, n_heads, n_units, rotary_base, rotary_scale, **_):
147
+ super().__init__()
148
+ self.attn = RotaryCrossAttention(idim, text_dim, n_units, n_heads, rotary_base, rotary_scale)
149
+ self.norm = ChannelLayerNorm(idim)
150
+
151
+ def forward(self, x, mask, text_emb, text_mask):
152
+ x = x * mask
153
+ y = self.attn(x.transpose(1, 2), text_emb.transpose(1, 2), mask.transpose(1, 2), text_mask.transpose(1, 2))
154
+ x = x + y.transpose(1, 2) * mask
155
+ return self.norm(x) * mask
156
+
157
+
158
+ class StyleCondBlock(nn.Module):
159
+ def __init__(self, idim, style_dim, n_units=256, n_heads=2):
160
+ super().__init__()
161
+ self.attention = StyleAttention(idim, style_dim, style_dim, n_units, n_heads, idim)
162
+ self.norm = ChannelLayerNorm(idim)
163
+
164
+ def forward(self, x, mask, style_key, style_value):
165
+ x = x * mask
166
+ m = mask.transpose(1, 2)
167
+ y = self.attention(x.transpose(1, 2), style_key, style_value, m) * m
168
+ x = x + y.transpose(1, 2)
169
+ return self.norm(x) * mask
170
+
171
+
172
+ class VectorField(nn.Module):
173
+ """Flow-matching velocity estimator on compressed latents (B, 144, T)."""
174
+
175
+ def __init__(self, cfg):
176
+ super().__init__()
177
+ p = cfg["proj_in"]
178
+ ldim = p["ldim"] * p["chunk_compress_factor"]
179
+ self.proj_in = PaddedConv1d(ldim, p["odim"], 1, bias=False)
180
+ self.time_encoder = TimeEncoder(cfg["time_encoder"]["time_dim"], cfg["time_encoder"]["hdim"])
181
+ mb = cfg["main_blocks"]
182
+ blocks = []
183
+ for _ in range(mb["n_blocks"]):
184
+ blocks += [
185
+ ConvNeXtStack(**mb["convnext_0"]),
186
+ TimeCondBlock(**mb["time_cond_layer"]),
187
+ ConvNeXtStack(**mb["convnext_1"]),
188
+ TextCondBlock(**mb["text_cond_layer"]),
189
+ ConvNeXtStack(**mb["convnext_2"]),
190
+ StyleCondBlock(**mb["style_cond_layer"]),
191
+ ]
192
+ self.main_blocks = nn.ModuleList(blocks)
193
+ self.last_convnext = ConvNeXtStack(**cfg["last_convnext"])
194
+ self.proj_out = PaddedConv1d(p["odim"], ldim, 1, bias=False)
195
+
196
+ def forward(self, x, t, text_emb, text_mask, style_key, style_value, latent_mask):
197
+ t_emb = self.time_encoder(t)
198
+ h = self.proj_in(x) * latent_mask
199
+ for blk in self.main_blocks:
200
+ if isinstance(blk, ConvNeXtStack):
201
+ h = blk(h, latent_mask)
202
+ elif isinstance(blk, TimeCondBlock):
203
+ h = blk(h, latent_mask, t_emb)
204
+ elif isinstance(blk, TextCondBlock):
205
+ h = blk(h, latent_mask, text_emb, text_mask)
206
+ else:
207
+ h = blk(h, latent_mask, style_key, style_value)
208
+ h = self.last_convnext(h, latent_mask)
209
+ return self.proj_out(h) * latent_mask
210
+
211
+
212
+ class TextToLatent(nn.Module):
213
+ def __init__(self, cfg):
214
+ super().__init__()
215
+ self.cfg = cfg
216
+ self.text_encoder = TextEncoder(cfg["text_encoder"])
217
+ self.style_encoder = StyleEncoder(cfg["style_encoder"])
218
+ s = cfg["speech_prompted_text_encoder"]
219
+ self.speech_prompted_text_encoder = SpeechPromptedTextEncoder(s["text_dim"], s["style_dim"], s["n_units"], s["n_heads"])
220
+ self.uncond_masker = UncondMasker(**cfg["uncond_masker"])
221
+ self.vector_field = VectorField(cfg["vector_field"])
222
+ self.sig_min = cfg["flow_matching"]["sig_min"]
223
+
224
+ @property
225
+ def style_key(self):
226
+ return self.style_encoder.style_token_layer.style_key
227
+
228
+ def encode_text(self, text_ids, text_mask, style_ttl):
229
+ """== text_encoder.onnx. Returns text_emb (B, 256, T)."""
230
+ key = self.style_key.expand(text_ids.shape[0], -1, -1)
231
+ x = self.text_encoder(text_ids, text_mask)
232
+ return self.speech_prompted_text_encoder(x, text_mask, key, style_ttl)
233
+
234
+ def velocity(self, x, t, text_emb, text_mask, style_ttl, latent_mask, uncond=False):
235
+ """Single (conditional or unconditional) velocity prediction."""
236
+ b = x.shape[0]
237
+ if uncond:
238
+ um = self.uncond_masker
239
+ text_emb = um.text_special_token.expand(b, -1, text_emb.shape[-1])
240
+ key = um.style_key_special_token.expand(b, -1, -1)
241
+ style_ttl = um.style_value_special_token.expand(b, -1, -1)
242
+ else:
243
+ key = self.style_key.expand(b, -1, -1)
244
+ return self.vector_field(x, t, text_emb, text_mask, key, style_ttl, latent_mask)
245
+
246
+ def cfg_velocity(self, x, t, text_emb, text_mask, style_ttl, latent_mask, cfg_scale=4.0):
247
+ """CFG as baked into vector_estimator.onnx: v_u + 4 (v_c - v_u) = 4 v_c - 3 v_u."""
248
+ xx = torch.cat([x, x])
249
+ tt = torch.cat([t, t])
250
+ te = torch.cat([text_emb, self.uncond_masker.text_special_token.expand_as(text_emb)])
251
+ tm = torch.cat([text_mask, text_mask])
252
+ lm = torch.cat([latent_mask, latent_mask])
253
+ b = x.shape[0]
254
+ key = torch.cat([self.style_key.expand(b, -1, -1), self.uncond_masker.style_key_special_token.expand(b, -1, -1)])
255
+ val = torch.cat([style_ttl, self.uncond_masker.style_value_special_token.expand(b, -1, -1)])
256
+ v = self.vector_field(xx, tt, te, tm, key, val, lm)
257
+ v_c, v_u = v[:b], v[b:] # (slicing instead of chunk keeps the batch dim dynamic in ONNX)
258
+ return v_u + cfg_scale * (v_c - v_u)
259
+
260
+ def euler_step(self, x, step, total_step, text_emb, text_mask, style_ttl, latent_mask, cfg_scale=4.0):
261
+ """== vector_estimator.onnx (one Euler step on t = step / total_step)."""
262
+ t = step / total_step
263
+ v = self.cfg_velocity(x, t, text_emb, text_mask, style_ttl, latent_mask, cfg_scale)
264
+ return (x + v / total_step.view(-1, 1, 1)) * latent_mask
265
+
266
+
267
+ # ==========================================================================
268
+ # Speech autoencoder
269
+ # ==========================================================================
270
+ class LatentDecoder(nn.Module):
271
+ """== vocoder.onnx decoder: causal ConvNeXt, head outputs 512 samples/frame."""
272
+
273
+ def __init__(self, cfg):
274
+ super().__init__()
275
+ h = cfg["hdim"]
276
+ self.embed = PaddedConv1d(cfg["idim"], h, cfg["ksz_init"], causal=True)
277
+ self.convnext = ConvNeXtStack(
278
+ h, cfg["ksz"], cfg["intermediate_dim"], cfg["num_layers"], cfg["dilation_lst"], causal=True, wrapped_dwconv=True
279
+ ).convnext
280
+ self.final_norm = nn.Module()
281
+ self.final_norm.norm = nn.BatchNorm1d(h, eps=1e-5)
282
+ hd = cfg["head"]
283
+ self.head = nn.Module()
284
+ self.head.layer1 = PaddedConv1d(hd["idim"], hd["hdim"], hd["ksz"], causal=True)
285
+ self.head.act = nn.PReLU(1)
286
+ self.head.layer2 = nn.Conv1d(hd["hdim"], hd["odim"], 1, bias=False)
287
+
288
+ def forward(self, z):
289
+ """z: raw latent (B, 24, T_f) -> wav (B, T_f * 512)."""
290
+ x = self.embed(z)
291
+ for blk in self.convnext:
292
+ x = blk(x)
293
+ x = self.final_norm.norm(x)
294
+ x = self.head.layer2(self.head.act(self.head.layer1(x))) # (B, 512, T_f)
295
+ return x.transpose(1, 2).reshape(x.shape[0], -1)
296
+
297
+
298
+ class SpecProcessor(nn.Module):
299
+ """log |STFT| (1025) ++ log mel (228) = 1253 features at hop 512.
300
+
301
+ Frames are left-aligned to the decoder: frame i covers samples ending at
302
+ 512*(i+1) (left pad n_fft - hop), so frame count == ceil(len / 512).
303
+ """
304
+
305
+ def __init__(self, n_fft, win_length, hop_length, n_mels, sample_rate, eps, **_):
306
+ super().__init__()
307
+ import torchaudio
308
+
309
+ self.n_fft, self.hop, self.win = n_fft, hop_length, win_length
310
+ self.eps = eps
311
+ self.register_buffer("window", torch.hann_window(win_length), persistent=False)
312
+ fb = torchaudio.functional.melscale_fbanks(n_fft // 2 + 1, 0.0, sample_rate / 2, n_mels, sample_rate, norm="slaney", mel_scale="slaney")
313
+ self.register_buffer("mel_fb", fb, persistent=False)
314
+
315
+ def forward(self, wav):
316
+ n = wav.shape[-1]
317
+ n_frames = (n + self.hop - 1) // self.hop
318
+ wav = F.pad(wav, (self.n_fft - self.hop, n_frames * self.hop - n))
319
+ spec = torch.stft(wav, self.n_fft, self.hop, self.win, self.window, center=False, return_complex=True).abs()
320
+ mel = torch.matmul(spec.transpose(1, 2), self.mel_fb).transpose(1, 2)
321
+ return torch.cat([torch.log(spec + self.eps), torch.log(mel + self.eps)], dim=1)
322
+
323
+
324
+ class LatentEncoder(nn.Module):
325
+ """wav @ 44.1 kHz -> raw latent (B, 24, T_f)."""
326
+
327
+ def __init__(self, cfg):
328
+ super().__init__()
329
+ self.spec_processor = SpecProcessor(**cfg["spec_processor"])
330
+ h = cfg["hdim"]
331
+ self.embed = PaddedConv1d(cfg["idim"], h, cfg["ksz_init"])
332
+ self.convnext = ConvNeXtStack(h, cfg["ksz"], cfg["intermediate_dim"], cfg["num_layers"], cfg["dilation_lst"], wrapped_dwconv=True).convnext
333
+ self.final_norm = ChannelLayerNorm(h)
334
+ self.proj_out = nn.Conv1d(h, cfg["odim"], 1)
335
+
336
+ def forward(self, wav):
337
+ x = self.embed(self.spec_processor(wav))
338
+ for blk in self.convnext:
339
+ x = blk(x)
340
+ return self.proj_out(self.final_norm(x))
341
+
342
+
343
+ class SpeechAutoencoder(nn.Module):
344
+ def __init__(self, cfg, ttl_cfg):
345
+ super().__init__()
346
+ self.sample_rate = cfg["sample_rate"]
347
+ self.hop = cfg["base_chunk_size"]
348
+ self.ccf = ttl_cfg["chunk_compress_factor"]
349
+ self.ldim = cfg["ldim"]
350
+ self.register_buffer("latent_mean", torch.zeros(1, cfg["ldim"], 1))
351
+ self.register_buffer("latent_std", torch.ones(1, cfg["ldim"], 1))
352
+ self.register_buffer("normalizer_scale", torch.tensor(float(ttl_cfg["normalizer"]["scale"])))
353
+ self.encoder = LatentEncoder(cfg["encoder"])
354
+ self.decoder = LatentDecoder(cfg["decoder"])
355
+
356
+ # ---- latent <-> TTL representation -----------------------------------
357
+ def compress(self, z):
358
+ """raw (B, 24, T_f) -> TTL latent (B, 144, ceil(T_f/6)); pads with the latent mean."""
359
+ b, c, t = z.shape
360
+ zn = (z - self.latent_mean) / self.latent_std
361
+ pad = (-t) % self.ccf
362
+ if pad:
363
+ zn = F.pad(zn, (0, pad))
364
+ zn = zn.view(b, c, -1, self.ccf).permute(0, 1, 3, 2).reshape(b, c * self.ccf, -1)
365
+ return zn * self.normalizer_scale
366
+
367
+ def decompress(self, x):
368
+ """TTL latent (B, 144, T) -> raw (B, 24, 6T)."""
369
+ b, _, t = x.shape
370
+ x = x / self.normalizer_scale
371
+ z = x.view(b, self.ldim, self.ccf, t).permute(0, 1, 3, 2).reshape(b, self.ldim, t * self.ccf)
372
+ return z * self.latent_std + self.latent_mean
373
+
374
+ def decode_ttl(self, x):
375
+ """== vocoder.onnx: TTL latent (B, 144, T) -> wav (B, 3072 T)."""
376
+ return self.decoder(self.decompress(x))
377
+
378
+ def encode_ttl(self, wav):
379
+ """wav (B, N) @ 44.1 kHz -> TTL latent (B, 144, ceil(N/3072))."""
380
+ return self.compress(self.encoder(wav))
381
+
382
+
383
+ # ==========================================================================
384
+ # Duration predictor
385
+ # ==========================================================================
386
+ class SentenceEncoder(nn.Module):
387
+ def __init__(self, cfg):
388
+ super().__init__()
389
+ d = cfg["char_emb_dim"]
390
+ self.sentence_token = nn.Parameter(torch.randn(1, d, 1) * 0.02)
391
+ self.text_embedder = CharEmbedder(N_VOCAB, cfg["text_embedder"]["char_emb_dim"])
392
+ self.convnext = ConvNeXtStack(**cfg["convnext"])
393
+ self.attn_encoder = AttnEncoder(**cfg["attn_encoder"])
394
+ self.proj_out = PaddedConv1d(cfg["proj_out"]["idim"], cfg["proj_out"]["odim"], 1, bias=False)
395
+
396
+ def forward(self, text_ids, text_mask):
397
+ b = text_ids.shape[0]
398
+ x = self.text_embedder(text_ids, text_mask)
399
+ x = torch.cat([self.sentence_token.expand(b, -1, -1), x], dim=-1)
400
+ m = torch.cat([torch.ones_like(text_mask[:, :, :1]), text_mask], dim=-1)
401
+ x = self.convnext(x, m)
402
+ x = self.attn_encoder(x, m) + x
403
+ return (self.proj_out(x[:, :, :1]) * m[:, :, :1]).flatten(1) # (B, 64)
404
+
405
+
406
+ class DurationHead(nn.Module):
407
+ def __init__(self, sentence_dim, n_style, style_dim, hdim, n_layer):
408
+ super().__init__()
409
+ assert n_layer == 2
410
+ self.layers = nn.ModuleList([nn.Linear(sentence_dim + n_style * style_dim, hdim), nn.Linear(hdim, 1)])
411
+ self.activation = nn.PReLU(1)
412
+
413
+ def forward(self, s, style_dp):
414
+ h = torch.cat([s, style_dp.flatten(1)], dim=1)
415
+ return self.layers[1](self.activation(self.layers[0](h))).squeeze(1) # log-seconds
416
+
417
+
418
+ class DurationPredictor(nn.Module):
419
+ def __init__(self, cfg):
420
+ super().__init__()
421
+ self.sentence_encoder = SentenceEncoder(cfg["sentence_encoder"])
422
+ self.style_encoder = StyleEncoder(cfg["style_encoder"])
423
+ self.predictor = DurationHead(**cfg["predictor"])
424
+
425
+ def log_duration(self, text_ids, text_mask, style_dp):
426
+ return self.predictor(self.sentence_encoder(text_ids, text_mask), style_dp)
427
+
428
+ def forward(self, text_ids, text_mask, style_dp):
429
+ """== duration_predictor.onnx: total utterance duration in seconds (B,)."""
430
+ return torch.exp(self.log_duration(text_ids, text_mask, style_dp))
431
+
432
+
433
+ # ==========================================================================
434
+ class LGTM(nn.Module):
435
+ def __init__(self, cfg):
436
+ super().__init__()
437
+ self.cfg = cfg
438
+ self.ttl = TextToLatent(cfg["ttl"])
439
+ self.ae = SpeechAutoencoder(cfg["ae"], cfg["ttl"])
440
+ self.dp = DurationPredictor(cfg["dp"])
441
+
lgtm/modules.py ADDED
@@ -0,0 +1,325 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """LGTM building blocks.
2
+
3
+ Tensor layout: channels-first (B, C, T); masks are (B, 1, T) float tensors of 0/1.
4
+ """
5
+ import math
6
+
7
+ import torch
8
+ import torch.nn as nn
9
+ import torch.nn.functional as F
10
+
11
+
12
+ # --------------------------------------------------------------------------
13
+ # Basic layers
14
+ # --------------------------------------------------------------------------
15
+ class ChannelLayerNorm(nn.Module):
16
+ """LayerNorm over channels of a (B, C, T) tensor. ONNX eps = 1e-6."""
17
+
18
+ def __init__(self, dim, eps=1e-6):
19
+ super().__init__()
20
+ self.norm = nn.LayerNorm(dim, eps=eps)
21
+
22
+ def forward(self, x):
23
+ return self.norm(x.transpose(1, 2)).transpose(1, 2)
24
+
25
+
26
+ class Linear(nn.Module):
27
+ """nn.Linear wrapped as ``.linear`` to match original names (``W_query.linear.weight``)."""
28
+
29
+ def __init__(self, idim, odim, bias=True):
30
+ super().__init__()
31
+ self.linear = nn.Linear(idim, odim, bias=bias)
32
+
33
+ def forward(self, x):
34
+ return self.linear(x)
35
+
36
+
37
+ class PaddedConv1d(nn.Module):
38
+ """Conv1d with replicate ('edge') padding, wrapped as ``.net`` like the original.
39
+
40
+ causal=True pads (k-1)*d on the left only (autoencoder decoder),
41
+ otherwise (k-1)*d/2 on both sides (text encoder / vector field).
42
+ """
43
+
44
+ def __init__(self, idim, odim, ksz, dilation=1, groups=1, bias=True, causal=False):
45
+ super().__init__()
46
+ self.net = nn.Conv1d(idim, odim, ksz, dilation=dilation, groups=groups, bias=bias)
47
+ total = (ksz - 1) * dilation
48
+ self.pad = (total, 0) if causal else (total // 2, total - total // 2)
49
+
50
+ def forward(self, x):
51
+ if self.pad != (0, 0):
52
+ x = F.pad(x, self.pad, mode="replicate")
53
+ return self.net(x)
54
+
55
+
56
+ class ConvNeXtBlock(nn.Module):
57
+ """ConvNeXt-1D block: dwconv -> LN -> pw(4x) -> GELU(erf) -> pw -> gamma, residual.
58
+
59
+ With a mask: input, dwconv output and block output are multiplied by the
60
+ mask (exactly as in the text encoder / vector-field graphs). The decoder
61
+ uses no mask and causal padding.
62
+ """
63
+
64
+ def __init__(self, dim, intermediate_dim, ksz, dilation=1, causal=False, wrapped_dwconv=False):
65
+ super().__init__()
66
+ self.gamma = nn.Parameter(torch.full((1, dim, 1), 1e-6))
67
+ dw = PaddedConv1d(dim, dim, ksz, dilation=dilation, groups=dim, causal=causal)
68
+ if wrapped_dwconv:
69
+ self.dwconv = dw # params: dwconv.net.{weight,bias}
70
+ else:
71
+ # params: dwconv.{weight,bias}; keep padding logic in the block.
72
+ self.dwconv = dw.net
73
+ self._pad = dw.pad
74
+ self.wrapped = wrapped_dwconv
75
+ self.norm = ChannelLayerNorm(dim)
76
+ self.pwconv1 = nn.Conv1d(dim, intermediate_dim, 1)
77
+ self.act = nn.GELU() # exact erf GELU, as in the graph
78
+ self.pwconv2 = nn.Conv1d(intermediate_dim, dim, 1)
79
+
80
+ def forward(self, x, mask=None):
81
+ if mask is not None:
82
+ x = x * mask
83
+ residual = x
84
+ if self.wrapped:
85
+ y = self.dwconv(x)
86
+ else:
87
+ y = self.dwconv(F.pad(x, self._pad, mode="replicate"))
88
+ if mask is not None:
89
+ y = y * mask
90
+ y = self.norm(y)
91
+ y = self.pwconv2(self.act(self.pwconv1(y)))
92
+ x = residual + self.gamma * y
93
+ if mask is not None:
94
+ x = x * mask
95
+ return x
96
+
97
+
98
+ class ConvNeXtStack(nn.Module):
99
+ """Stack of ConvNeXt blocks, params at ``convnext.{i}.*``."""
100
+
101
+ def __init__(self, idim, ksz, intermediate_dim, num_layers, dilation_lst, causal=False, wrapped_dwconv=False, **_):
102
+ super().__init__()
103
+ assert len(dilation_lst) == num_layers
104
+ self.convnext = nn.ModuleList(
105
+ [
106
+ ConvNeXtBlock(idim, intermediate_dim, ksz, d, causal=causal, wrapped_dwconv=wrapped_dwconv)
107
+ for d in dilation_lst
108
+ ]
109
+ )
110
+
111
+ def forward(self, x, mask=None):
112
+ for blk in self.convnext:
113
+ x = blk(x, mask)
114
+ return x
115
+
116
+
117
+ # --------------------------------------------------------------------------
118
+ # VITS-style relative-position self-attention encoder (text / sentence enc.)
119
+ # --------------------------------------------------------------------------
120
+ class RelPosMultiHeadAttention(nn.Module):
121
+ """VITS MultiHeadAttention with shared relative position embeddings (window 4)."""
122
+
123
+ def __init__(self, channels, n_heads, window_size=4):
124
+ super().__init__()
125
+ assert channels % n_heads == 0
126
+ self.n_heads = n_heads
127
+ self.k_channels = channels // n_heads
128
+ self.window_size = window_size
129
+ self.conv_q = nn.Conv1d(channels, channels, 1)
130
+ self.conv_k = nn.Conv1d(channels, channels, 1)
131
+ self.conv_v = nn.Conv1d(channels, channels, 1)
132
+ self.conv_o = nn.Conv1d(channels, channels, 1)
133
+ std = self.k_channels ** -0.5
134
+ self.emb_rel_k = nn.Parameter(torch.randn(1, 2 * window_size + 1, self.k_channels) * std)
135
+ self.emb_rel_v = nn.Parameter(torch.randn(1, 2 * window_size + 1, self.k_channels) * std)
136
+
137
+ def forward(self, x, attn_mask):
138
+ q, k, v = self.conv_q(x), self.conv_k(x), self.conv_v(x)
139
+ b, d, t = k.shape
140
+ h, kc = self.n_heads, self.k_channels
141
+ query = q.view(b, h, kc, t).transpose(2, 3) / math.sqrt(kc)
142
+ key = k.view(b, h, kc, t).transpose(2, 3)
143
+ value = v.view(b, h, kc, t).transpose(2, 3)
144
+
145
+ scores = torch.matmul(query, key.transpose(-2, -1))
146
+ key_rel = self._get_relative_embeddings(self.emb_rel_k, t)
147
+ rel_logits = torch.matmul(query, key_rel.unsqueeze(0).transpose(-2, -1))
148
+ scores = scores + self._relative_to_absolute(rel_logits)
149
+ scores = scores.masked_fill(attn_mask == 0, -1e4)
150
+ p = F.softmax(scores, dim=-1)
151
+ out = torch.matmul(p, value)
152
+ value_rel = self._get_relative_embeddings(self.emb_rel_v, t)
153
+ out = out + torch.matmul(self._absolute_to_relative(p), value_rel.unsqueeze(0))
154
+ out = out.transpose(2, 3).contiguous().view(b, d, t)
155
+ return self.conv_o(out)
156
+
157
+ def _get_relative_embeddings(self, emb, length):
158
+ w = self.window_size
159
+ pad_length = max(length - (w + 1), 0)
160
+ start = max((w + 1) - length, 0)
161
+ emb = F.pad(emb, (0, 0, pad_length, pad_length, 0, 0)) # pad of 0 is a no-op (keeps export branch-free)
162
+ return emb[:, start: start + 2 * length - 1]
163
+
164
+ @staticmethod
165
+ def _relative_to_absolute(x):
166
+ b, h, l, _ = x.shape
167
+ x = F.pad(x, (0, 1))
168
+ x = x.view(b, h, l * 2 * l)
169
+ x = F.pad(x, (0, l - 1))
170
+ return x.view(b, h, l + 1, 2 * l - 1)[:, :, :l, l - 1:]
171
+
172
+ @staticmethod
173
+ def _absolute_to_relative(x):
174
+ b, h, l, _ = x.shape
175
+ x = F.pad(x, (0, l - 1))
176
+ x = x.view(b, h, l * l + l * (l - 1))
177
+ x = F.pad(x, (l, 0))
178
+ return x.view(b, h, l, 2 * l)[:, :, :, 1:]
179
+
180
+
181
+ class FFN(nn.Module):
182
+ def __init__(self, channels, filter_channels):
183
+ super().__init__()
184
+ self.conv_1 = nn.Conv1d(channels, filter_channels, 1)
185
+ self.conv_2 = nn.Conv1d(filter_channels, channels, 1)
186
+
187
+ def forward(self, x, mask):
188
+ x = torch.relu(self.conv_1(x * mask))
189
+ return self.conv_2(x * mask) * mask
190
+
191
+
192
+ class AttnEncoder(nn.Module):
193
+ """VITS post-norm transformer encoder."""
194
+
195
+ def __init__(self, hidden_channels, filter_channels, n_heads, n_layers, p_dropout=0.0):
196
+ super().__init__()
197
+ self.attn_layers = nn.ModuleList([RelPosMultiHeadAttention(hidden_channels, n_heads) for _ in range(n_layers)])
198
+ self.norm_layers_1 = nn.ModuleList([ChannelLayerNorm(hidden_channels) for _ in range(n_layers)])
199
+ self.ffn_layers = nn.ModuleList([FFN(hidden_channels, filter_channels) for _ in range(n_layers)])
200
+ self.norm_layers_2 = nn.ModuleList([ChannelLayerNorm(hidden_channels) for _ in range(n_layers)])
201
+ self.drop = nn.Dropout(p_dropout)
202
+
203
+ def forward(self, x, mask):
204
+ attn_mask = mask.unsqueeze(2) * mask.unsqueeze(-1)
205
+ x = x * mask
206
+ for attn, n1, ffn, n2 in zip(self.attn_layers, self.norm_layers_1, self.ffn_layers, self.norm_layers_2):
207
+ x = n1(x + self.drop(attn(x, attn_mask)))
208
+ x = n2(x + self.drop(ffn(x, mask)))
209
+ return x * mask
210
+
211
+
212
+ class CharEmbedder(nn.Module):
213
+ def __init__(self, n_vocab, dim):
214
+ super().__init__()
215
+ self.char_embedder = nn.Embedding(n_vocab, dim)
216
+
217
+ def forward(self, ids, mask):
218
+ return self.char_embedder(ids).transpose(1, 2) * mask
219
+
220
+
221
+ # --------------------------------------------------------------------------
222
+ # Style (GST-like) cross attention: keys go through tanh.
223
+ # --------------------------------------------------------------------------
224
+ class StyleAttention(nn.Module):
225
+ """Multi-head cross-attention to style tokens (used in the text encoder and
226
+ the vector field's style-conditioning layers).
227
+
228
+ Heads are formed by splitting the last dim and stacking on a new leading
229
+ axis; keys are passed through tanh; scores are divided by ``sqrt(n_units)``;
230
+ rows of padded queries are zeroed after softmax.
231
+ """
232
+
233
+ def __init__(self, q_dim, k_dim, v_dim, n_units, n_heads, out_dim):
234
+ super().__init__()
235
+ self.n_heads = n_heads
236
+ self.n_units = n_units
237
+ self.W_query = Linear(q_dim, n_units)
238
+ self.W_key = Linear(k_dim, n_units)
239
+ self.W_value = Linear(v_dim, n_units)
240
+ self.out_fc = Linear(n_units, out_dim)
241
+
242
+ def _heads(self, x): # (B, T, U) -> (H, B, T, U/H)
243
+ return torch.stack(torch.chunk(x, self.n_heads, dim=-1), dim=0)
244
+
245
+ def forward(self, q, k, v, q_mask=None):
246
+ """q: (B, Tq, q_dim), k: (B, Tk, k_dim), v: (B, Tk, v_dim), q_mask: (B, Tq, 1)."""
247
+ q = self._heads(self.W_query(q))
248
+ k = torch.tanh(self._heads(self.W_key(k)))
249
+ v = self._heads(self.W_value(v))
250
+ p = F.softmax(torch.matmul(q, k.transpose(-1, -2)) / math.sqrt(self.n_units), dim=-1)
251
+ if q_mask is not None:
252
+ p = p * q_mask.unsqueeze(0)
253
+ out = torch.matmul(p, v) # (H, B, Tq, d)
254
+ out = torch.cat(out.unbind(0), dim=-1)
255
+ return self.out_fc(out)
256
+
257
+
258
+ # --------------------------------------------------------------------------
259
+ # Rotary cross attention (latent queries -> text keys), vector field.
260
+ # --------------------------------------------------------------------------
261
+ class RotaryCrossAttention(nn.Module):
262
+ """Text-conditioning attention in the vector field.
263
+
264
+ Rotary embedding uses *length-normalised* positions: pos = i / length,
265
+ angle = pos * theta, theta_j = rotary_scale * base^(-j/(d/2)). Queries use
266
+ the latent length, keys the text length (so the attention learns a soft
267
+ monotonic alignment in relative position). Non-interleaved rotation
268
+ (first half / second half). Scores are divided by sqrt(n_units / 2).
269
+ """
270
+
271
+ def __init__(self, idim, text_dim, n_units, n_heads, rotary_base=10000, rotary_scale=10, **_):
272
+ super().__init__()
273
+ self.n_heads = n_heads
274
+ self.head_dim = n_units // n_heads
275
+ self.scale = math.sqrt(n_units / 2) # = 16 for n_units=512 (verified against ONNX)
276
+ self.W_query = Linear(idim, n_units)
277
+ self.W_key = Linear(text_dim, n_units)
278
+ self.W_value = Linear(text_dim, n_units)
279
+ self.out_fc = Linear(n_units, idim)
280
+ half = self.head_dim // 2
281
+ theta = rotary_scale * rotary_base ** (-torch.arange(half, dtype=torch.float32) / half)
282
+ self.register_buffer("theta", theta.view(1, 1, half))
283
+
284
+ def _heads(self, x): # (B, T, U) -> (H, B, T, d)
285
+ b, t, _ = x.shape
286
+ return x.view(b, t, self.n_heads, self.head_dim).permute(2, 0, 1, 3)
287
+
288
+ def _rotate(self, x, mask):
289
+ # x: (H, B, T, d); mask: (B, T, 1)
290
+ t = x.shape[2]
291
+ length = mask.sum(dim=(1, 2)).view(-1, 1, 1)
292
+ pos = torch.arange(t, device=x.device, dtype=x.dtype).view(1, t, 1) / length
293
+ ang = pos * self.theta # (B, T, d/2)
294
+ cos, sin = torch.cos(ang), torch.sin(ang)
295
+ x1, x2 = x.chunk(2, dim=-1)
296
+ return torch.cat([x1 * cos - x2 * sin, x1 * sin + x2 * cos], dim=-1)
297
+
298
+ def forward(self, x, text, x_mask, text_mask):
299
+ """x: (B, Tq, idim), text: (B, Tk, text_dim), masks (B, T, 1)."""
300
+ q = self._rotate(self._heads(self.W_query(x)), x_mask)
301
+ k = self._rotate(self._heads(self.W_key(text)), text_mask)
302
+ v = self._heads(self.W_value(text))
303
+ scores = torch.matmul(q, k.transpose(-1, -2)) / self.scale
304
+ key_mask = text_mask.transpose(1, 2).unsqueeze(0) # (1, B, 1, Tk)
305
+ scores = scores.masked_fill(key_mask == 0, float("-inf"))
306
+ p = F.softmax(scores, dim=-1) * x_mask.unsqueeze(0)
307
+ out = torch.matmul(p, v) # (H, B, Tq, d)
308
+ b, tq = out.shape[1], out.shape[2]
309
+ out = out.permute(1, 2, 0, 3).reshape(b, tq, -1)
310
+ return self.out_fc(out)
311
+
312
+
313
+ class TimeEncoder(nn.Module):
314
+ """sinusoidal(t * 1000) -> Linear -> Mish -> Linear."""
315
+
316
+ def __init__(self, time_dim, hdim):
317
+ super().__init__()
318
+ half = time_dim // 2
319
+ freqs = torch.exp(-math.log(10000) * torch.arange(half, dtype=torch.float32) / (half - 1))
320
+ self.register_buffer("freqs", freqs.view(1, half))
321
+ self.mlp = nn.Sequential(Linear(time_dim, hdim), nn.Mish(), Linear(hdim, time_dim))
322
+
323
+ def forward(self, t): # t: (B,) in [0, 1]
324
+ ang = t.view(-1, 1) * 1000.0 * self.freqs
325
+ return self.mlp(torch.cat([torch.sin(ang), torch.cos(ang)], dim=-1))
lgtm/onnx_inference.py ADDED
@@ -0,0 +1,150 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """ONNX Runtime inference for LGTM (numpy + onnxruntime only).
2
+
3
+ from lgtm import LGTMOnnx
4
+ tts = LGTMOnnx.from_pretrained("path/or/hf-repo", use_gpu=False)
5
+ wav = tts.synthesize("Bonjour à tous.", lang="fr", voice="M1")
6
+ tts.save_wav(wav, "out.wav")
7
+ voice = tts.clone_voice("reference.wav")
8
+ """
9
+ import json
10
+ import os
11
+
12
+ import numpy as np
13
+ import onnxruntime as ort
14
+ import soundfile as sf
15
+ from scipy.signal import resample_poly
16
+
17
+ from .text import AVAILABLE_LANGS, TextProcessor
18
+
19
+ SAMPLE_RATE = 44100
20
+ CHUNK = 512 * 6 # samples per latent frame
21
+
22
+
23
+ def _split_text(text, max_len):
24
+ import re
25
+
26
+ chunks = []
27
+ for para in [p.strip() for p in re.split(r"\n\s*\n+", text.strip()) if p.strip()]:
28
+ cur = ""
29
+ for sent in re.split(r"(?<=[.!?。!?])\s+", para):
30
+ if cur and len(cur) + len(sent) + 1 > max_len:
31
+ chunks.append(cur)
32
+ cur = sent
33
+ else:
34
+ cur = f"{cur} {sent}".strip()
35
+ if cur:
36
+ chunks.append(cur)
37
+ return chunks
38
+
39
+
40
+ def _prepare_reference(w, top_db=40.0, pad_s=0.1, target_rms_db=-23.0, peak=0.95):
41
+ hop = SAMPLE_RATE // 100
42
+ n = len(w) // hop
43
+ if n >= 3:
44
+ db = 10 * np.log10((w[: n * hop].reshape(n, hop) ** 2).mean(1) + 1e-10)
45
+ active = np.nonzero(db > db.max() - top_db)[0]
46
+ if len(active):
47
+ pad = int(pad_s * SAMPLE_RATE)
48
+ w = w[max(active[0] * hop - pad, 0): min((active[-1] + 1) * hop + pad, len(w))]
49
+ w = w * (10 ** (target_rms_db / 20) / (np.sqrt((w ** 2).mean()) + 1e-6))
50
+ m = np.abs(w).max()
51
+ return (w * (peak / m) if m > peak else w).astype(np.float32)
52
+
53
+
54
+ def load_voice_style(path):
55
+ d = json.load(open(path))
56
+ ttl = np.array(d["style_ttl"]["data"], np.float32).reshape(d["style_ttl"]["dims"])
57
+ dp = np.array(d["style_dp"]["data"], np.float32).reshape(d["style_dp"]["dims"])
58
+ return ttl, dp
59
+
60
+
61
+ def save_voice_style(path, voice):
62
+ ttl, dp = voice
63
+ json.dump({"style_ttl": {"data": ttl.tolist(), "dims": list(ttl.shape), "type": "float32"},
64
+ "style_dp": {"data": dp.tolist(), "dims": list(dp.shape), "type": "float32"}}, open(path, "w"))
65
+
66
+
67
+ class LGTMOnnx:
68
+ def __init__(self, model_dir, use_gpu=False, threads=None):
69
+ providers = (["CUDAExecutionProvider"] if use_gpu else []) + ["CPUExecutionProvider"]
70
+ so = ort.SessionOptions()
71
+ if threads:
72
+ so.intra_op_num_threads = threads
73
+ load = lambda n: ort.InferenceSession(os.path.join(model_dir, "onnx", f"{n}.onnx"), so, providers=providers) # noqa: E731
74
+ self.text_encoder = load("text_encoder")
75
+ self.duration_predictor = load("duration_predictor")
76
+ self.vector_estimator = load("vector_estimator")
77
+ self.vocoder = load("vocoder")
78
+ self._voice_encoder_path = os.path.join(model_dir, "onnx", "voice_encoder.onnx")
79
+ self._voice_encoder = None
80
+ self._so, self._providers = so, providers
81
+ self.tp = TextProcessor(os.path.join(model_dir, "unicode_indexer.json"))
82
+ self.voice_dir = os.path.join(model_dir, "voice_styles")
83
+
84
+ @classmethod
85
+ def from_pretrained(cls, repo_or_dir, **kw):
86
+ if not os.path.isdir(repo_or_dir):
87
+ from huggingface_hub import snapshot_download
88
+
89
+ repo_or_dir = snapshot_download(repo_or_dir, allow_patterns=["unicode_indexer.json", "voice_styles/*", "onnx/*"])
90
+ return cls(repo_or_dir, **kw)
91
+
92
+ @property
93
+ def voices(self):
94
+ return sorted(f[:-5] for f in os.listdir(self.voice_dir) if f.endswith(".json"))
95
+
96
+ def _voice(self, voice):
97
+ if isinstance(voice, str):
98
+ return load_voice_style(voice if voice.endswith(".json") else os.path.join(self.voice_dir, f"{voice}.json"))
99
+ return voice
100
+
101
+ def clone_voice(self, wav_path, max_seconds=15.0):
102
+ """Voice style from a reference recording (5-15 s of clean speech works best)."""
103
+ if self._voice_encoder is None:
104
+ self._voice_encoder = ort.InferenceSession(self._voice_encoder_path, self._so, providers=self._providers)
105
+ w, sr = sf.read(wav_path, dtype="float32")
106
+ if w.ndim == 2:
107
+ w = w.mean(1)
108
+ if sr != SAMPLE_RATE:
109
+ g = np.gcd(sr, SAMPLE_RATE)
110
+ w = resample_poly(w, SAMPLE_RATE // g, sr // g).astype(np.float32)
111
+ w = _prepare_reference(w)[: int(max_seconds * SAMPLE_RATE)]
112
+ ttl, dp = self._voice_encoder.run(None, {"wav": w[None]})
113
+ return ttl, dp
114
+
115
+ def _synth_batch(self, texts, langs, voice, steps, speed, rng):
116
+ s_ttl, s_dp = voice
117
+ b = len(texts)
118
+ ids, mask = self.tp.batch_numpy(texts, langs)
119
+ s_ttl = np.repeat(s_ttl, b, axis=0) if s_ttl.shape[0] == 1 else s_ttl
120
+ s_dp = np.repeat(s_dp, b, axis=0) if s_dp.shape[0] == 1 else s_dp
121
+ dur = self.duration_predictor.run(None, {"text_ids": ids, "style_dp": s_dp, "text_mask": mask})[0] / speed
122
+ text_emb = self.text_encoder.run(None, {"text_ids": ids, "style_ttl": s_ttl, "text_mask": mask})[0]
123
+ wav_len = (dur * SAMPLE_RATE).astype(np.int64)
124
+ lat_len = (wav_len + CHUNK - 1) // CHUNK
125
+ L = int(lat_len.max())
126
+ lmask = (np.arange(L)[None] < lat_len[:, None]).astype(np.float32)[:, None]
127
+ x = rng.standard_normal((b, 144, L)).astype(np.float32) * lmask
128
+ total = np.full(b, steps, np.float32)
129
+ for i in range(steps):
130
+ x = self.vector_estimator.run(None, {"noisy_latent": x, "text_emb": text_emb, "style_ttl": s_ttl, "latent_mask": lmask,
131
+ "text_mask": mask, "current_step": np.full(b, i, np.float32), "total_step": total})[0]
132
+ wav = self.vocoder.run(None, {"latent": x})[0]
133
+ return [wav[k, : wav_len[k]] for k in range(b)]
134
+
135
+ def synthesize(self, text, lang="en", voice="F1", steps=8, speed=1.05, silence=0.3, seed=None):
136
+ """Text -> 44.1 kHz float32 waveform. voice: preset name, .json path, or clone_voice() output."""
137
+ if lang not in AVAILABLE_LANGS:
138
+ raise ValueError(f"unsupported language {lang!r}")
139
+ rng = np.random.default_rng(seed)
140
+ chunks = _split_text(text, 120 if lang in ("ja", "ko") else 300)
141
+ wavs = self._synth_batch(chunks, [lang] * len(chunks), self._voice(voice), steps, speed, rng)
142
+ gap = np.zeros(int(silence * SAMPLE_RATE), np.float32)
143
+ out = []
144
+ for i, w in enumerate(wavs):
145
+ out += [w] + ([gap] if i < len(wavs) - 1 else [])
146
+ return np.concatenate(out)
147
+
148
+ @staticmethod
149
+ def save_wav(wav, path):
150
+ sf.write(path, wav, SAMPLE_RATE)
lgtm/strip_onnx.py ADDED
@@ -0,0 +1,43 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Remove exporter debug metadata (stack traces with local file paths, doc strings) from ONNX files."""
2
+ import sys
3
+
4
+ import onnx
5
+
6
+
7
+ def _clean_graph(g):
8
+ g.doc_string = ""
9
+ del g.metadata_props[:]
10
+ for vi in list(g.input) + list(g.output) + list(g.value_info):
11
+ vi.doc_string = ""
12
+ del vi.metadata_props[:]
13
+ for init in g.initializer:
14
+ init.doc_string = ""
15
+ del init.metadata_props[:]
16
+ for n in g.node:
17
+ n.doc_string = ""
18
+ del n.metadata_props[:]
19
+ for a in n.attribute:
20
+ if a.g is not None and a.HasField("g"):
21
+ _clean_graph(a.g)
22
+ for sg in a.graphs:
23
+ _clean_graph(sg)
24
+
25
+
26
+ def strip(path):
27
+ m = onnx.load(path)
28
+ m.doc_string = ""
29
+ del m.metadata_props[:]
30
+ _clean_graph(m.graph)
31
+ for f in m.functions:
32
+ f.doc_string = ""
33
+ del f.metadata_props[:]
34
+ for n in f.node:
35
+ n.doc_string = ""
36
+ del n.metadata_props[:]
37
+ onnx.save(m, path)
38
+
39
+
40
+ if __name__ == "__main__":
41
+ for p in sys.argv[1:]:
42
+ strip(p)
43
+ print("stripped", p)
lgtm/text.py ADDED
@@ -0,0 +1,81 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Text frontend: NFKD normalisation, light cleanup, "<lang>text</lang>" tagging, and a
2
+ code-point -> id lookup (unicode_indexer.json, 8322 ids)."""
3
+ import json
4
+ import re
5
+ from unicodedata import normalize
6
+
7
+
8
+ AVAILABLE_LANGS = ["en", "ko", "ja", "ar", "bg", "cs", "da", "de", "el", "es", "et", "fi", "fr", "hi", "hr", "hu", "id", "it", "lt", "lv", "nl", "pl", "pt", "ro", "ru", "sk", "sl", "sv", "tr", "uk", "vi", "na"]
9
+
10
+ _EMOJI = re.compile(
11
+ "[\U0001f600-\U0001f64f\U0001f300-\U0001f5ff\U0001f680-\U0001f6ff\U0001f700-\U0001f77f"
12
+ "\U0001f780-\U0001f7ff\U0001f800-\U0001f8ff\U0001f900-\U0001f9ff\U0001fa00-\U0001fa6f"
13
+ "\U0001fa70-\U0001faff☀-⛿✀-➿\U0001f1e6-\U0001f1ff]+",
14
+ flags=re.UNICODE,
15
+ )
16
+ _REPL = {"–": "-", "‑": "-", "—": "-", "_": " ", "“": '"', "”": '"', "‘": "'", "’": "'",
17
+ "´": "'", "`": "'", "[": " ", "]": " ", "|": " ", "/": " ", "#": " ", "→": " ", "←": " "}
18
+ _EXPR = {"@": " at ", "e.g.,": "for example, ", "i.e.,": "that is, "}
19
+
20
+
21
+ def preprocess_text(text: str, lang: str) -> str:
22
+ text = normalize("NFKD", text)
23
+ text = _EMOJI.sub("", text)
24
+ for k, v in _REPL.items():
25
+ text = text.replace(k, v)
26
+ text = re.sub(r"[♥☆♡©\\]", "", text)
27
+ for k, v in _EXPR.items():
28
+ text = text.replace(k, v)
29
+ for p in [",", r"\.", "!", r"\?", ";", ":", "'"]:
30
+ text = re.sub(r" " + p, p.replace("\\", ""), text)
31
+ while '""' in text:
32
+ text = text.replace('""', '"')
33
+ while "''" in text:
34
+ text = text.replace("''", "'")
35
+ while "``" in text:
36
+ text = text.replace("``", "`")
37
+ text = re.sub(r"\s+", " ", text).strip()
38
+ if not re.search(r"[.!?;:,'\"')\]}…。」』】〉》›»]$", text):
39
+ text += "."
40
+ if lang not in AVAILABLE_LANGS:
41
+ raise ValueError(f"Invalid language: {lang}")
42
+ return f"<{lang}>" + text + f"</{lang}>"
43
+
44
+
45
+ class TextProcessor:
46
+ def __init__(self, unicode_indexer_path: str):
47
+ self.indexer_path = unicode_indexer_path
48
+ with open(unicode_indexer_path) as f:
49
+ self.indexer = json.load(f) # list: code point -> id (-1 = unknown)
50
+
51
+ def encode(self, text: str, lang: str, preprocess: bool = True) -> list[int]:
52
+ s = preprocess_text(text, lang) if preprocess else text
53
+ return [self.indexer[ord(c)] if ord(c) < len(self.indexer) else -1 for c in s]
54
+
55
+ def unknown_chars(self, text: str, lang: str) -> set[str]:
56
+ s = preprocess_text(text, lang)
57
+ return {c for c in s if ord(c) >= len(self.indexer) or self.indexer[ord(c)] < 0}
58
+
59
+ def batch_numpy(self, texts, langs):
60
+ """ids (B, T) int64 and mask (B, 1, T) float32 as numpy arrays."""
61
+ import numpy as np
62
+
63
+ seqs = [self.encode(t, l) for t, l in zip(texts, langs)]
64
+ T = max(len(s) for s in seqs)
65
+ ids = np.zeros((len(seqs), T), np.int64)
66
+ mask = np.zeros((len(seqs), 1, T), np.float32)
67
+ for i, s in enumerate(seqs):
68
+ ids[i, : len(s)] = s
69
+ mask[i, :, : len(s)] = 1
70
+ return ids, mask
71
+
72
+ def batch(self, texts: list[str], langs: list[str], device=None):
73
+ import torch
74
+
75
+ seqs = [self.encode(t, l) for t, l in zip(texts, langs)]
76
+ lengths = torch.tensor([len(s) for s in seqs])
77
+ ids = torch.zeros(len(seqs), int(lengths.max()), dtype=torch.long)
78
+ for i, s in enumerate(seqs):
79
+ ids[i, : len(s)] = torch.tensor(s)
80
+ mask = (torch.arange(ids.shape[1])[None] < lengths[:, None]).float().unsqueeze(1)
81
+ return ids.to(device), mask.to(device)
onnx/duration_predictor.onnx ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:6fbdf338c3768273d92ab1d14ed80f1db7214ac49e2377d1e81d6240bf8a8486
3
+ size 3512757
onnx/text_encoder.onnx ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:58e20a45e3a52d6755b35ffff1501fd990ae4cfd8f7431cf5ca561e422efe69b
3
+ size 36088450
onnx/vector_estimator.onnx ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:829a06f43e524bf082658b5dd5d97f7e8e7082fa551fa80b829faf8127e64ba2
3
+ size 256228310
onnx/vocoder.onnx ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:96c75b381f26e701df5b9849304890310a9ebe7d8f585da8e5ec60bf3f8e8567
3
+ size 101387918
onnx/voice_encoder.onnx ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:594b3c8710a7e2f73ddbd1122500f6219a1aa7dd5394cb6f22dacc454f219a71
3
+ size 117859547
pytorch/model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:dcb429cd3bee3f37574bb47abaf0f9154de31adcc2f8570a248f011c59cf8c2e
3
+ size 513767008
requirements-onnx.txt ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ onnxruntime>=1.17 # or onnxruntime-gpu
2
+ numpy
3
+ soundfile
4
+ scipy
5
+ huggingface_hub
requirements.txt ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ torch>=2.1
2
+ torchaudio>=2.1
3
+ numpy
4
+ safetensors
5
+ soundfile
6
+ huggingface_hub
unicode_indexer.json ADDED
The diff for this file is too large to render. See raw diff
 
voice_styles/F1.json ADDED
The diff for this file is too large to render. See raw diff
 
voice_styles/F2.json ADDED
The diff for this file is too large to render. See raw diff
 
voice_styles/F3.json ADDED
The diff for this file is too large to render. See raw diff
 
voice_styles/F4.json ADDED
The diff for this file is too large to render. See raw diff
 
voice_styles/F5.json ADDED
The diff for this file is too large to render. See raw diff
 
voice_styles/M1.json ADDED
The diff for this file is too large to render. See raw diff
 
voice_styles/M2.json ADDED
The diff for this file is too large to render. See raw diff
 
voice_styles/M3.json ADDED
The diff for this file is too large to render. See raw diff
 
voice_styles/M4.json ADDED
The diff for this file is too large to render. See raw diff
 
voice_styles/M5.json ADDED
The diff for this file is too large to render. See raw diff