ema-lightning 1.0.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.
- ema_lightning/__init__.py +5 -0
- ema_lightning/api.py +247 -0
- ema_lightning/audio.py +67 -0
- ema_lightning/chunker.py +44 -0
- ema_lightning/decoder.py +65 -0
- ema_lightning/engine.py +158 -0
- ema_lightning/frontend.py +59 -0
- ema_lightning/graphs.py +123 -0
- ema_lightning/model.py +247 -0
- ema_lightning/scheduler.py +169 -0
- ema_lightning-1.0.0.dist-info/METADATA +329 -0
- ema_lightning-1.0.0.dist-info/RECORD +14 -0
- ema_lightning-1.0.0.dist-info/WHEEL +4 -0
- ema_lightning-1.0.0.dist-info/licenses/LICENSE +201 -0
ema_lightning/api.py
ADDED
|
@@ -0,0 +1,247 @@
|
|
|
1
|
+
"""EMA Lightning: Turkish text to speech.
|
|
2
|
+
|
|
3
|
+
tts = EMA().lightning()
|
|
4
|
+
speech = tts.say("Merhaba, nasılsınız?", path="merhaba.wav")
|
|
5
|
+
speeches = tts.say(["Birinci cümle.", "İkinci cümle."])
|
|
6
|
+
for chunk in tts.stream("Uzun bir metin..."):
|
|
7
|
+
play(chunk)
|
|
8
|
+
"""
|
|
9
|
+
import json
|
|
10
|
+
import os
|
|
11
|
+
import random
|
|
12
|
+
import statistics
|
|
13
|
+
import time
|
|
14
|
+
import warnings
|
|
15
|
+
from dataclasses import dataclass, field
|
|
16
|
+
from pathlib import Path
|
|
17
|
+
|
|
18
|
+
import numpy as np
|
|
19
|
+
import torch
|
|
20
|
+
from huggingface_hub import hf_hub_download
|
|
21
|
+
|
|
22
|
+
from .audio import RATE, RATES, Resampler, write_wav
|
|
23
|
+
from .chunker import chunk
|
|
24
|
+
from .decoder import load_decoder
|
|
25
|
+
from .engine import FIRST_WINDOW, Engine, windows
|
|
26
|
+
from .frontend import Frontend
|
|
27
|
+
from .graphs import Graphs
|
|
28
|
+
from .model import load_acoustic
|
|
29
|
+
from .scheduler import DONE, Playhead
|
|
30
|
+
|
|
31
|
+
REPO = "canberkkkkkk/ema-lightning"
|
|
32
|
+
PROBE = "Bugün hava çok güzel, yarın da yağmur yağacakmış; toplantı öğleden sonra başlayacak."
|
|
33
|
+
CACHE = Path(os.environ.get("XDG_CACHE_HOME", Path.home() / ".cache")) / "ema_lightning" / "batch_size_v2.json"
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
@dataclass(frozen=True)
|
|
37
|
+
class Speech:
|
|
38
|
+
"""One spoken text: float32 audio in [-1, 1] at sample_rate, and the seed that made it."""
|
|
39
|
+
|
|
40
|
+
audio: np.ndarray = field(repr=False)
|
|
41
|
+
sample_rate: int
|
|
42
|
+
duration: float
|
|
43
|
+
seed: int
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class EMA:
|
|
47
|
+
def __init__(self, device="auto"):
|
|
48
|
+
device = torch.device(("cuda" if torch.cuda.is_available() else "cpu") if device == "auto" else device)
|
|
49
|
+
model = load_acoustic(hf_hub_download(REPO, "ema.pt"), device)
|
|
50
|
+
decoder = load_decoder(hf_hub_download(REPO, "decoder.pt"), device)
|
|
51
|
+
self._setup(model, decoder, Frontend(model.vocab), device)
|
|
52
|
+
|
|
53
|
+
@classmethod
|
|
54
|
+
def _from_parts(cls, model, decoder, frontend, device):
|
|
55
|
+
self = cls.__new__(cls)
|
|
56
|
+
self._setup(model, decoder, frontend, device)
|
|
57
|
+
return self
|
|
58
|
+
|
|
59
|
+
def _setup(self, model, decoder, frontend, device):
|
|
60
|
+
self.device = torch.device(device)
|
|
61
|
+
self._engine = Engine(model, decoder, self.device)
|
|
62
|
+
self._frontend = frontend
|
|
63
|
+
self._batch_size = None
|
|
64
|
+
self._playhead = Playhead(self._engine, self._list_batch)
|
|
65
|
+
|
|
66
|
+
def lightning(self, batch_size=None):
|
|
67
|
+
"""Compile and record every stage as CUDA graphs, check them against the plain path, report ready."""
|
|
68
|
+
if self.device.type != "cuda":
|
|
69
|
+
warnings.warn("lightning needs a CUDA GPU; EMA keeps working without it", stacklevel=2)
|
|
70
|
+
return self
|
|
71
|
+
if self._engine.graphs is not None:
|
|
72
|
+
return self
|
|
73
|
+
# cuDNN times each convolution shape once and keeps the fastest method. The decoder is all convolutions:
|
|
74
|
+
# on an RTX PRO 6000 this made it about twice as fast. It is process-wide and changes speed, not the audio.
|
|
75
|
+
torch.backends.cudnn.benchmark = True
|
|
76
|
+
start = time.perf_counter()
|
|
77
|
+
size = batch_size or self.best_batch_size()
|
|
78
|
+
try:
|
|
79
|
+
graphs = Graphs(self._engine.model, self._engine.decoder, size)
|
|
80
|
+
except Exception as error: # e.g. no compiler on this machine: graphs alone still help
|
|
81
|
+
warnings.warn(f"compiling failed, recording graphs without it: {error}", stacklevel=2)
|
|
82
|
+
graphs = Graphs(self._engine.model, self._engine.decoder, size, compile=False)
|
|
83
|
+
built = time.perf_counter() - start
|
|
84
|
+
plain = self._probe()
|
|
85
|
+
self._engine.graphs = graphs
|
|
86
|
+
fast = self._probe()
|
|
87
|
+
if any((a - b).abs().max() > 1e-2 for a, b in zip(plain, fast, strict=True)):
|
|
88
|
+
self._engine.graphs = None
|
|
89
|
+
warnings.warn("recorded graphs disagree with the plain path; lightning stays off", stacklevel=2)
|
|
90
|
+
return self
|
|
91
|
+
self._batch_size = size
|
|
92
|
+
first = statistics.median(self._first_audio() for _ in range(5))
|
|
93
|
+
timings = [self._timed_probe() for _ in range(5)]
|
|
94
|
+
ms, seconds = statistics.median(t for t, _ in timings), timings[0][1]
|
|
95
|
+
print(f"EMA Lightning ready: {graphs.count} graphs built in {built:.1f} s, first audio in {first:.1f} ms, "
|
|
96
|
+
f"a {seconds:.1f} s sentence in {ms:.1f} ms, batch size {size}")
|
|
97
|
+
return self
|
|
98
|
+
|
|
99
|
+
def best_batch_size(self):
|
|
100
|
+
"""The smallest batch that reaches 90% of this device's best throughput, measured once and cached."""
|
|
101
|
+
if self._batch_size:
|
|
102
|
+
return self._batch_size
|
|
103
|
+
key = torch.cuda.get_device_name(self.device) if self.device.type == "cuda" else f"cpu-{os.cpu_count()}"
|
|
104
|
+
cached = json.loads(CACHE.read_text()) if CACHE.exists() else {}
|
|
105
|
+
if key not in cached:
|
|
106
|
+
sizes = (1, 2, 4, 8, 16, 32, 64, 128) if self.device.type == "cuda" else (1, 2, 4, 8)
|
|
107
|
+
rates = {}
|
|
108
|
+
for size in sizes:
|
|
109
|
+
try:
|
|
110
|
+
rates[size] = self._throughput(size)
|
|
111
|
+
except torch.cuda.OutOfMemoryError:
|
|
112
|
+
torch.cuda.empty_cache()
|
|
113
|
+
break
|
|
114
|
+
cached[key] = min(s for s, r in rates.items() if r >= 0.9 * max(rates.values()))
|
|
115
|
+
CACHE.parent.mkdir(parents=True, exist_ok=True)
|
|
116
|
+
CACHE.write_text(json.dumps(cached))
|
|
117
|
+
self._batch_size = cached[key]
|
|
118
|
+
return self._batch_size
|
|
119
|
+
|
|
120
|
+
def say(self, text, speed=1.0, seed=None, sample_rate=RATE, path=None):
|
|
121
|
+
"""Speech for one text, or a list of Speech for a list of texts (batched)."""
|
|
122
|
+
_check(speed, seed, sample_rate)
|
|
123
|
+
if isinstance(text, str):
|
|
124
|
+
seed = _seed(seed)
|
|
125
|
+
request = self._submit(text, speed, seed)
|
|
126
|
+
speech = _speech(list(self._receive(request, sample_rate)), sample_rate, seed)
|
|
127
|
+
if path is not None:
|
|
128
|
+
write_wav(path, speech.audio, sample_rate)
|
|
129
|
+
return speech
|
|
130
|
+
texts = _texts(text)
|
|
131
|
+
seeds = [_seed(seed) for _ in texts]
|
|
132
|
+
requests = [self._submit(t, speed, s) for t, s in zip(texts, seeds, strict=True)]
|
|
133
|
+
out = [_speech(list(self._receive(r, sample_rate)), sample_rate, s)
|
|
134
|
+
for r, s in zip(requests, seeds, strict=True)]
|
|
135
|
+
if path is not None:
|
|
136
|
+
Path(path).mkdir(parents=True, exist_ok=True)
|
|
137
|
+
width = len(str(max(len(out) - 1, 0)))
|
|
138
|
+
for i, speech in enumerate(out):
|
|
139
|
+
write_wav(Path(path) / f"{i:0{width}d}.wav", speech.audio, sample_rate)
|
|
140
|
+
return out
|
|
141
|
+
|
|
142
|
+
def stream(self, text, speed=1.0, seed=None, sample_rate=RATE):
|
|
143
|
+
"""float32 audio chunks as they are made: one second first, then four seconds at a time.
|
|
144
|
+
|
|
145
|
+
Call it from as many threads as you like; every stream shares the GPU through the same queues.
|
|
146
|
+
Stopping early (leaving the loop) drops the rest of that text's work.
|
|
147
|
+
"""
|
|
148
|
+
_check(speed, seed, sample_rate)
|
|
149
|
+
if not isinstance(text, str):
|
|
150
|
+
raise TypeError("stream() takes one text; for many, call stream() once per text")
|
|
151
|
+
return self._stream(text, speed, _seed(seed), sample_rate)
|
|
152
|
+
|
|
153
|
+
def _stream(self, text, speed, seed, sample_rate):
|
|
154
|
+
"""A stream's request, submitted when its first chunk is asked for."""
|
|
155
|
+
yield from self._receive(self._submit(text, speed, seed, first=FIRST_WINDOW), sample_rate)
|
|
156
|
+
|
|
157
|
+
def _submit(self, text, speed, seed, first=None):
|
|
158
|
+
pieces = self._pieces(text, speed, seed)
|
|
159
|
+
return self._playhead.submit(pieces, speed) if first is None else self._playhead.submit(pieces, speed, first)
|
|
160
|
+
|
|
161
|
+
def _receive(self, request, sample_rate):
|
|
162
|
+
"""float32 chunks from a request's outbox, resampled. If the caller stops waiting, its work is dropped."""
|
|
163
|
+
resampler = Resampler(sample_rate)
|
|
164
|
+
try:
|
|
165
|
+
while (item := request.outbox.get()) is not DONE:
|
|
166
|
+
if isinstance(item, BaseException):
|
|
167
|
+
raise item
|
|
168
|
+
out = resampler.push(item)
|
|
169
|
+
if out.numel():
|
|
170
|
+
yield out.float().cpu().numpy()
|
|
171
|
+
tail = resampler.flush()
|
|
172
|
+
if tail is not None and tail.numel():
|
|
173
|
+
yield tail.float().cpu().numpy()
|
|
174
|
+
finally:
|
|
175
|
+
request.cancel()
|
|
176
|
+
|
|
177
|
+
def _pieces(self, text, speed, seed):
|
|
178
|
+
return [self._engine.piece(piece, pause, (seed * 1_000_003 + i) % 2**63)
|
|
179
|
+
for i, (piece, pause) in enumerate(chunk(self._frontend(text), speed))]
|
|
180
|
+
|
|
181
|
+
def _list_batch(self):
|
|
182
|
+
return self._batch_size or self.best_batch_size()
|
|
183
|
+
|
|
184
|
+
def _probe(self):
|
|
185
|
+
"""Latents and both window sizes of a fixed sentence, for checking graphs against the plain path."""
|
|
186
|
+
(piece,) = self._pieces(PROBE, 1.0, 0)
|
|
187
|
+
with self._engine.lock:
|
|
188
|
+
self._engine.plan([piece], 1.0)
|
|
189
|
+
self._engine.think([piece])
|
|
190
|
+
short = self._engine.decode([(piece, windows(piece.frames, FIRST_WINDOW)[0])])[0]
|
|
191
|
+
full = self._engine.decode([(piece, windows(piece.frames)[0])])[0]
|
|
192
|
+
return piece.latents, short, full
|
|
193
|
+
|
|
194
|
+
def _first_audio(self):
|
|
195
|
+
"""Milliseconds until a stream of the probe sentence hands over its first chunk."""
|
|
196
|
+
start = time.perf_counter()
|
|
197
|
+
next(iter(self._stream(PROBE, 1.0, 0, RATE)))
|
|
198
|
+
return 1000 * (time.perf_counter() - start)
|
|
199
|
+
|
|
200
|
+
def _timed_probe(self):
|
|
201
|
+
"""Milliseconds for say() to return the probe sentence, and the seconds of audio it made."""
|
|
202
|
+
start = time.perf_counter()
|
|
203
|
+
speech = self.say(PROBE, seed=0)
|
|
204
|
+
return 1000 * (time.perf_counter() - start), speech.duration
|
|
205
|
+
|
|
206
|
+
def _throughput(self, size):
|
|
207
|
+
"""Seconds of audio per second for a batch of `size` copies of the probe sentence."""
|
|
208
|
+
(piece,) = self._pieces(PROBE, 1.0, 0)
|
|
209
|
+
pieces = [self._engine.piece(piece.text, 0.0, i) for i in range(size)]
|
|
210
|
+
best = float("inf")
|
|
211
|
+
with self._engine.lock:
|
|
212
|
+
self._engine.plan(pieces, 1.0)
|
|
213
|
+
for _ in range(3):
|
|
214
|
+
if self.device.type == "cuda":
|
|
215
|
+
torch.cuda.synchronize(self.device)
|
|
216
|
+
start = time.perf_counter()
|
|
217
|
+
self._engine.think(pieces)
|
|
218
|
+
for span in windows(pieces[0].frames):
|
|
219
|
+
self._engine.decode([(p, span) for p in pieces])
|
|
220
|
+
if self.device.type == "cuda":
|
|
221
|
+
torch.cuda.synchronize(self.device)
|
|
222
|
+
best = min(best, time.perf_counter() - start)
|
|
223
|
+
return size * pieces[0].frames / 25 / best
|
|
224
|
+
|
|
225
|
+
|
|
226
|
+
def _check(speed, seed, sample_rate):
|
|
227
|
+
if isinstance(speed, bool) or not isinstance(speed, (int, float)) or not 0.25 <= speed <= 4:
|
|
228
|
+
raise ValueError("speed must be a number from 0.25 to 4")
|
|
229
|
+
if seed is not None and (isinstance(seed, bool) or not isinstance(seed, int) or seed < 0):
|
|
230
|
+
raise ValueError("seed must be a non-negative integer")
|
|
231
|
+
if sample_rate not in RATES:
|
|
232
|
+
raise ValueError(f"sample_rate must be one of {RATES}")
|
|
233
|
+
|
|
234
|
+
|
|
235
|
+
def _texts(text):
|
|
236
|
+
if not isinstance(text, (list, tuple)) or not all(isinstance(t, str) for t in text):
|
|
237
|
+
raise TypeError("text must be a string or a list of strings")
|
|
238
|
+
return list(text)
|
|
239
|
+
|
|
240
|
+
|
|
241
|
+
def _seed(seed):
|
|
242
|
+
return random.SystemRandom().randrange(2**31) if seed is None else seed
|
|
243
|
+
|
|
244
|
+
|
|
245
|
+
def _speech(chunks, sample_rate, seed):
|
|
246
|
+
audio = np.concatenate(chunks).astype(np.float32) if chunks else np.zeros(0, np.float32)
|
|
247
|
+
return Speech(audio, sample_rate, len(audio) / sample_rate, seed)
|
ema_lightning/audio.py
ADDED
|
@@ -0,0 +1,67 @@
|
|
|
1
|
+
"""Audio helpers: sample-rate conversion from 48 kHz, window by window, and WAV writing.
|
|
2
|
+
|
|
3
|
+
The resampler carries its state from window to window, so the output never depends on where the decoder's
|
|
4
|
+
windows were cut.
|
|
5
|
+
"""
|
|
6
|
+
import wave
|
|
7
|
+
|
|
8
|
+
import numpy as np
|
|
9
|
+
import torch
|
|
10
|
+
|
|
11
|
+
RATE = 48000
|
|
12
|
+
RATES = (48000, 24000, 16000, 8000)
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class Resampler:
|
|
16
|
+
"""48 kHz in, 48 kHz / factor out, through a linear-phase windowed-sinc low-pass."""
|
|
17
|
+
|
|
18
|
+
def __init__(self, rate):
|
|
19
|
+
self.factor = RATE // rate
|
|
20
|
+
self.half = 32 * self.factor
|
|
21
|
+
self.taps = None
|
|
22
|
+
self.buffer = None
|
|
23
|
+
self.start = -self.half # absolute input index of buffer[0]; the signal is preceded by zeros
|
|
24
|
+
self.next = 0 # absolute input index at the centre of the next output sample
|
|
25
|
+
self.end = 0 # absolute input index one past the last real sample
|
|
26
|
+
|
|
27
|
+
def push(self, chunk):
|
|
28
|
+
if self.factor == 1:
|
|
29
|
+
return chunk
|
|
30
|
+
if self.buffer is None:
|
|
31
|
+
self.buffer = chunk.new_zeros(self.half)
|
|
32
|
+
n = torch.arange(-self.half, self.half + 1, dtype=torch.float64)
|
|
33
|
+
cutoff = 0.45 / self.factor
|
|
34
|
+
taps = 2 * cutoff * torch.sinc(2 * cutoff * n) * torch.blackman_window(2 * self.half + 1, False,
|
|
35
|
+
dtype=torch.float64)
|
|
36
|
+
self.taps = (taps / taps.sum()).to(chunk)[None, None]
|
|
37
|
+
self.buffer = torch.cat([self.buffer, chunk])
|
|
38
|
+
self.end += chunk.numel()
|
|
39
|
+
return self._emit(self.end - 1 - self.half)
|
|
40
|
+
|
|
41
|
+
def flush(self):
|
|
42
|
+
if self.factor == 1 or self.buffer is None:
|
|
43
|
+
return None
|
|
44
|
+
self.buffer = torch.cat([self.buffer, self.buffer.new_zeros(self.half + self.factor)])
|
|
45
|
+
return self._emit(self.end - 1)
|
|
46
|
+
|
|
47
|
+
def _emit(self, last):
|
|
48
|
+
"""Every output whose centre is at or before input index `last`."""
|
|
49
|
+
if last < self.next:
|
|
50
|
+
return self.buffer.new_zeros(0)
|
|
51
|
+
count = (last - self.next) // self.factor + 1
|
|
52
|
+
a = self.next - self.half - self.start
|
|
53
|
+
b = a + (count - 1) * self.factor + 2 * self.half + 1
|
|
54
|
+
out = torch.nn.functional.conv1d(self.buffer[a:b][None, None], self.taps, stride=self.factor)[0, 0]
|
|
55
|
+
self.next += count * self.factor
|
|
56
|
+
drop = self.next - self.half - self.start
|
|
57
|
+
self.buffer, self.start = self.buffer[drop:], self.start + drop
|
|
58
|
+
return out
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def write_wav(path, audio, rate):
|
|
62
|
+
pcm = (np.clip(audio, -1.0, 1.0) * 32767.0).round().astype("<i2")
|
|
63
|
+
with wave.open(str(path), "wb") as f:
|
|
64
|
+
f.setnchannels(1)
|
|
65
|
+
f.setsampwidth(2)
|
|
66
|
+
f.setframerate(rate)
|
|
67
|
+
f.writeframes(pcm.tobytes())
|
ema_lightning/chunker.py
ADDED
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
"""Greedy chunking: spoken text in, pieces short enough for one pass of the model out.
|
|
2
|
+
|
|
3
|
+
Text that fits in about ten seconds of speech stays one piece; the model reads its punctuation itself.
|
|
4
|
+
Longer text is cut at the last good spot inside each ten-second window: a sentence end, then a clause
|
|
5
|
+
mark, then a space, and only when there is none of those, exactly at the limit.
|
|
6
|
+
"""
|
|
7
|
+
import re
|
|
8
|
+
|
|
9
|
+
LETTERS_PER_SECOND = 18.0
|
|
10
|
+
MAX_SECONDS = 10.0
|
|
11
|
+
MAX_LETTERS = 250
|
|
12
|
+
SENTENCE_PAUSE = 0.25
|
|
13
|
+
CLAUSE_PAUSE = 0.12
|
|
14
|
+
CUTS = ((re.compile(r"[.!?]+[\"')]*(?= )"), SENTENCE_PAUSE), (re.compile(r"[,;:](?= )"), CLAUSE_PAUSE),
|
|
15
|
+
(re.compile(r"\S(?= )"), CLAUSE_PAUSE))
|
|
16
|
+
LETTER = re.compile(r"[^\W\d_]")
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def chunk(text, speed):
|
|
20
|
+
"""[(piece, seconds of silence after it)], each piece ending in terminal punctuation."""
|
|
21
|
+
limit = int(min(MAX_LETTERS, LETTERS_PER_SECOND * MAX_SECONDS * speed))
|
|
22
|
+
pieces, rest = [], text.strip()
|
|
23
|
+
while rest:
|
|
24
|
+
cut, pause = len(rest), 0.0
|
|
25
|
+
if len(rest) > limit:
|
|
26
|
+
cut = limit
|
|
27
|
+
for pattern, gap in CUTS:
|
|
28
|
+
ends = [m.end() for m in pattern.finditer(rest, 0, limit + 1)]
|
|
29
|
+
if ends:
|
|
30
|
+
cut, pause = ends[-1], gap
|
|
31
|
+
break
|
|
32
|
+
piece, rest = rest[:cut].strip(), rest[cut:].strip()
|
|
33
|
+
if LETTER.search(piece):
|
|
34
|
+
pieces.append((finish(piece), pause))
|
|
35
|
+
if pieces:
|
|
36
|
+
pieces[-1] = (pieces[-1][0], 0.0)
|
|
37
|
+
return pieces
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def finish(piece):
|
|
41
|
+
"""End every piece the way the model was trained: on a sentence end."""
|
|
42
|
+
if piece.rstrip("\"')")[-1:] in (".", "!", "?"):
|
|
43
|
+
return piece
|
|
44
|
+
return piece.rstrip(",;:- ") + "."
|
ema_lightning/decoder.py
ADDED
|
@@ -0,0 +1,65 @@
|
|
|
1
|
+
"""EMA Lightning's decoder: 64-dim latents at 25 Hz in, 48 kHz audio out.
|
|
2
|
+
|
|
3
|
+
With `lengths`, everything past each row's real length is held at zero after every layer, so a window
|
|
4
|
+
padded to a fixed shape decodes exactly like the same window unpadded.
|
|
5
|
+
"""
|
|
6
|
+
import math
|
|
7
|
+
|
|
8
|
+
import torch
|
|
9
|
+
import torch.nn as nn
|
|
10
|
+
import torch.nn.functional as F
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def _keep(x, length):
|
|
14
|
+
return x if length is None else x * (torch.arange(x.shape[-1], device=x.device) < length[:, None])[:, None]
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class ResBlock(nn.Module):
|
|
18
|
+
def __init__(self, ch, k, dilations):
|
|
19
|
+
super().__init__()
|
|
20
|
+
self.convs1 = nn.ModuleList(nn.Conv1d(ch, ch, k, dilation=d, padding=(k * d - d) // 2) for d in dilations)
|
|
21
|
+
self.convs2 = nn.ModuleList(nn.Conv1d(ch, ch, k, padding=(k - 1) // 2) for _ in dilations)
|
|
22
|
+
|
|
23
|
+
def forward(self, x, length=None):
|
|
24
|
+
for c1, c2 in zip(self.convs1, self.convs2, strict=True):
|
|
25
|
+
y = _keep(F.leaky_relu(c1(F.leaky_relu(x, 0.1)), 0.1), length)
|
|
26
|
+
x = _keep(x + c2(y), length)
|
|
27
|
+
return x
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
class Decoder(nn.Module):
|
|
31
|
+
def __init__(self, latent_dim=64, ch=256, rates=(8, 6, 5, 2, 2, 2), kernels=(16, 12, 10, 4, 4, 4),
|
|
32
|
+
rb_kernels=(3, 5, 9), rb_dilations=((1, 3, 5),) * 3):
|
|
33
|
+
super().__init__()
|
|
34
|
+
self.hop, self.nk = math.prod(rates), len(rb_kernels)
|
|
35
|
+
self.pre = nn.Conv1d(latent_dim, ch, 7, padding=3)
|
|
36
|
+
self.ups = nn.ModuleList(nn.ConvTranspose1d(ch >> i, ch >> (i + 1), k, r, padding=(k - r) // 2)
|
|
37
|
+
for i, (r, k) in enumerate(zip(rates, kernels, strict=True)))
|
|
38
|
+
self.blocks = nn.ModuleList(ResBlock(ch >> (i + 1), k, d)
|
|
39
|
+
for i in range(len(rates)) for k, d in zip(rb_kernels, rb_dilations, strict=True))
|
|
40
|
+
self.post = nn.Conv1d(ch >> len(rates), 1, 7, padding=3)
|
|
41
|
+
|
|
42
|
+
def forward(self, z, lengths=None):
|
|
43
|
+
"""z: [B, latent_dim, T]; lengths: real frames per row, or None when nothing is padded."""
|
|
44
|
+
frames = z.shape[-1]
|
|
45
|
+
length = lengths
|
|
46
|
+
x = _keep(self.pre(_keep(z, length)), length)
|
|
47
|
+
for i, up in enumerate(self.ups):
|
|
48
|
+
x = up(F.leaky_relu(x, 0.1))
|
|
49
|
+
if length is not None:
|
|
50
|
+
length = (length - 1) * up.stride[0] - 2 * up.padding[0] + up.kernel_size[0]
|
|
51
|
+
x = _keep(x, length)
|
|
52
|
+
x = sum(b(x, length) for b in self.blocks[i * self.nk:(i + 1) * self.nk]) / self.nk
|
|
53
|
+
return torch.tanh(self.post(F.leaky_relu(x)))[:, 0, :frames * self.hop]
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def load_decoder(path, device):
|
|
57
|
+
ck = torch.load(path, map_location="cpu", weights_only=False)
|
|
58
|
+
cfg, sd = ck.get("cfg", {}), ck["G"]
|
|
59
|
+
for key in [k for k in sd if k.endswith("weight_g")]:
|
|
60
|
+
g, v = sd.pop(key), sd.pop(key[:-1] + "v")
|
|
61
|
+
sd[key[:-2]] = g * v / v.norm(dim=tuple(range(1, v.dim())), keepdim=True)
|
|
62
|
+
keys = ("latent_dim", "ch", "rates", "kernels", "rb_kernels", "rb_dilations")
|
|
63
|
+
model = Decoder(**{k: cfg[k] for k in keys if k in cfg})
|
|
64
|
+
model.load_state_dict(sd)
|
|
65
|
+
return model.to(device).float().eval().requires_grad_(False)
|
ema_lightning/engine.py
ADDED
|
@@ -0,0 +1,158 @@
|
|
|
1
|
+
"""The speech engine: pieces of spoken text in, 48 kHz audio out, one decoded window at a time.
|
|
2
|
+
|
|
3
|
+
Every piece goes through three stages: plan (the text stage and its frame timeline), think (the
|
|
4
|
+
aligner and the four steps) and decode (one window of latents to audio). Each stage takes any batch;
|
|
5
|
+
with lightning on, the same stages replay from recorded CUDA graphs instead. Playhead (scheduler.py)
|
|
6
|
+
decides which work runs in each batch.
|
|
7
|
+
"""
|
|
8
|
+
import threading
|
|
9
|
+
from dataclasses import dataclass
|
|
10
|
+
from itertools import pairwise
|
|
11
|
+
|
|
12
|
+
import torch
|
|
13
|
+
|
|
14
|
+
RATE = 48000
|
|
15
|
+
FIRST_WINDOW = 25 # a stream's first window is one second, so its first audio comes fast
|
|
16
|
+
WINDOW = 100 # four seconds for every later window
|
|
17
|
+
CONTEXT = 8 # frames decoded on each side of a window; the decoder reaches 4
|
|
18
|
+
# Every batch of windows is padded to one of these lengths before decoding. The GPU picks its convolution
|
|
19
|
+
# method by input length, and on an RTX PRO 6000 the windows' own lengths (33, 108, 116 frames) ran about
|
|
20
|
+
# twice as slow per frame as 48 and 120. The decoder ignores padding exactly, so the audio does not change.
|
|
21
|
+
DECODE_SIZES = (48, 120)
|
|
22
|
+
MAX_WORD_FRAMES = 250
|
|
23
|
+
MAX_FRAMES = 3000
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
@dataclass(eq=False)
|
|
27
|
+
class Piece:
|
|
28
|
+
text: str
|
|
29
|
+
pause: float
|
|
30
|
+
seed: int
|
|
31
|
+
ids: torch.Tensor
|
|
32
|
+
cw: torch.Tensor # word of each letter
|
|
33
|
+
wstart: torch.Tensor # first letter of that word
|
|
34
|
+
h: torch.Tensor = None
|
|
35
|
+
dur: torch.Tensor = None
|
|
36
|
+
fw: torch.Tensor = None # word of each frame
|
|
37
|
+
fp: torch.Tensor = None # position of each frame inside its word
|
|
38
|
+
latents: torch.Tensor = None
|
|
39
|
+
spans: list = None # its decoder windows, fixed when it is planned
|
|
40
|
+
|
|
41
|
+
@property
|
|
42
|
+
def letters(self):
|
|
43
|
+
return self.ids.numel()
|
|
44
|
+
|
|
45
|
+
@property
|
|
46
|
+
def frames(self):
|
|
47
|
+
return self.fw.numel()
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def windows(frames, first=WINDOW):
|
|
51
|
+
"""(start, end) of every decoded window: `first` frames first, then four seconds each."""
|
|
52
|
+
spans, s = [], 0
|
|
53
|
+
while s < frames:
|
|
54
|
+
e = min(frames, s + (first if s == 0 else WINDOW))
|
|
55
|
+
spans.append((s, e))
|
|
56
|
+
s = e
|
|
57
|
+
return spans
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def batches(items, size):
|
|
61
|
+
for i in range(0, len(items), size):
|
|
62
|
+
yield items[i:i + size]
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
class Engine:
|
|
66
|
+
def __init__(self, model, decoder, device):
|
|
67
|
+
self.model, self.decoder, self.device = model, decoder, torch.device(device)
|
|
68
|
+
self.hop = decoder.hop
|
|
69
|
+
self.graphs = None
|
|
70
|
+
self.lock = threading.RLock()
|
|
71
|
+
self.decode_sizes = DECODE_SIZES if self.device.type == "cuda" else () # slow lengths are a GPU problem
|
|
72
|
+
|
|
73
|
+
def piece(self, text, pause, seed):
|
|
74
|
+
ids = [self.model.stoi.get(ch, 1) for ch in text]
|
|
75
|
+
starts = [i for i, ch in enumerate(text) if ch != " " and (i == 0 or text[i - 1] == " ")] or [0]
|
|
76
|
+
bounds = [0] + starts[1:] + [len(text)]
|
|
77
|
+
cw, wstart = [], []
|
|
78
|
+
for w, (a, b) in enumerate(pairwise(bounds)):
|
|
79
|
+
cw += [w] * (b - a)
|
|
80
|
+
wstart += [a] * (b - a)
|
|
81
|
+
return Piece(text, pause, seed, torch.tensor(ids), torch.tensor(cw), torch.tensor(wstart))
|
|
82
|
+
|
|
83
|
+
@torch.no_grad()
|
|
84
|
+
def plan(self, pieces, speed):
|
|
85
|
+
"""Durations for a batch, then each piece's exact frame timeline after a single wait."""
|
|
86
|
+
L = max(p.letters for p in pieces)
|
|
87
|
+
ids = self.stack([p.ids for p in pieces], L, 0)
|
|
88
|
+
mask = ids != 0
|
|
89
|
+
h, dur = self.run("text", ids=ids, mask=mask)
|
|
90
|
+
dur = dur / speed
|
|
91
|
+
word = self.stack([p.cw for p in pieces], L, 0)
|
|
92
|
+
counts = torch.zeros_like(dur).scatter_add_(1, word, dur).round().clamp(1, MAX_WORD_FRAMES).long().cpu()
|
|
93
|
+
for i, p in enumerate(pieces):
|
|
94
|
+
n = counts[i, :int(p.cw[-1]) + 1]
|
|
95
|
+
frames = min(int(n.sum()), MAX_FRAMES)
|
|
96
|
+
fw = torch.repeat_interleave(torch.arange(n.numel()), n)[:frames]
|
|
97
|
+
fp = ((torch.arange(frames) - (n.cumsum(0) - n)[fw]).double() / n[fw].double()).float()
|
|
98
|
+
p.h, p.dur, p.fw, p.fp = h[i, :p.letters].clone(), dur[i, :p.letters].clone(), fw, fp
|
|
99
|
+
|
|
100
|
+
@torch.no_grad()
|
|
101
|
+
def think(self, pieces):
|
|
102
|
+
"""Latents for a batch, each piece from its own seeded noise."""
|
|
103
|
+
L, T = max(p.letters for p in pieces), max(p.frames for p in pieces)
|
|
104
|
+
latents = self.run(
|
|
105
|
+
"sound",
|
|
106
|
+
h=self.stack([p.h for p in pieces], L, 0.0),
|
|
107
|
+
dur=self.stack([p.dur for p in pieces], L, 0.0),
|
|
108
|
+
mask=self.stack([torch.ones(p.letters, dtype=torch.bool) for p in pieces], L, False),
|
|
109
|
+
cw=self.stack([p.cw for p in pieces], L, -1),
|
|
110
|
+
wstart=self.stack([p.wstart for p in pieces], L, 0),
|
|
111
|
+
fw=self.stack([p.fw for p in pieces], T, -1),
|
|
112
|
+
fp=self.stack([p.fp for p in pieces], T, 0.0),
|
|
113
|
+
fmask=self.stack([torch.ones(p.frames, dtype=torch.bool) for p in pieces], T, False),
|
|
114
|
+
noise=self.stack([self.noise(p) for p in pieces], T, 0.0, dim=1))
|
|
115
|
+
for i, p in enumerate(pieces):
|
|
116
|
+
p.latents = latents[i, :p.frames].clone()
|
|
117
|
+
|
|
118
|
+
@torch.no_grad()
|
|
119
|
+
def decode_size(self, piece, span):
|
|
120
|
+
"""The length a window is decoded at: itself plus its margins, padded to a fast size on a GPU."""
|
|
121
|
+
s, e = span
|
|
122
|
+
n = min(piece.frames, e + CONTEXT) - max(0, s - CONTEXT)
|
|
123
|
+
return next((size for size in self.decode_sizes if size >= n), n)
|
|
124
|
+
|
|
125
|
+
def decode(self, items):
|
|
126
|
+
"""Audio for a batch of (piece, window), each cut to its own window."""
|
|
127
|
+
spans = [(max(0, s - CONTEXT), min(p.frames, e + CONTEXT)) for p, (s, e) in items]
|
|
128
|
+
longest = max(b - a for a, b in spans)
|
|
129
|
+
P = next((size for size in self.decode_sizes if size >= longest), longest)
|
|
130
|
+
z = self.stack([p.latents[a:b] for (p, _), (a, b) in zip(items, spans, strict=True)], P, 0.0)
|
|
131
|
+
lengths = torch.tensor([b - a for a, b in spans], device=self.device)
|
|
132
|
+
audio = self.run("decode", z=z, lengths=lengths)
|
|
133
|
+
return [audio[i, (s - a) * self.hop:(e - a) * self.hop].clone()
|
|
134
|
+
for i, ((_, (s, e)), (a, _)) in enumerate(zip(items, spans, strict=True))]
|
|
135
|
+
|
|
136
|
+
def run(self, stage, **args):
|
|
137
|
+
out = self.graphs.run(stage, **args) if self.graphs is not None else None
|
|
138
|
+
if out is not None:
|
|
139
|
+
return out
|
|
140
|
+
if stage == "text":
|
|
141
|
+
return self.model.text_stage(**args)
|
|
142
|
+
if stage == "sound":
|
|
143
|
+
return self.model.sound_stage(**args)
|
|
144
|
+
full = bool((args["lengths"] == args["z"].shape[1]).all()) # nothing padded: skip the masks
|
|
145
|
+
return self.decoder(args["z"].transpose(1, 2), None if full else args["lengths"])
|
|
146
|
+
|
|
147
|
+
def noise(self, piece):
|
|
148
|
+
generator = torch.Generator(device=self.device).manual_seed(piece.seed)
|
|
149
|
+
shape = (len(self.model.times), piece.frames, self.model.latent_dim)
|
|
150
|
+
return torch.randn(shape, generator=generator, device=self.device)
|
|
151
|
+
|
|
152
|
+
def stack(self, rows, size, fill, dim=0):
|
|
153
|
+
shape = list(rows[0].shape)
|
|
154
|
+
shape[dim] = size
|
|
155
|
+
out = torch.full([len(rows)] + shape, fill, dtype=rows[0].dtype, device=self.device)
|
|
156
|
+
for i, row in enumerate(rows):
|
|
157
|
+
out[i].narrow(dim, 0, row.shape[dim]).copy_(row)
|
|
158
|
+
return out
|
|
@@ -0,0 +1,59 @@
|
|
|
1
|
+
"""Text frontend: any written Turkish in, text in the model's own alphabet out.
|
|
2
|
+
|
|
3
|
+
normalizer-tr reads numbers, dates, times, money, units, symbols and abbreviations aloud. Its
|
|
4
|
+
"fallback" policy leaves nothing unread: notation it cannot resolve is spoken literally instead of
|
|
5
|
+
being kept as written. The alphabet step then lowercases the Turkish way and drops whatever the
|
|
6
|
+
model cannot read. Nothing here ever raises on text.
|
|
7
|
+
"""
|
|
8
|
+
import re
|
|
9
|
+
import unicodedata
|
|
10
|
+
|
|
11
|
+
from normalizer_tr import Normalizer
|
|
12
|
+
|
|
13
|
+
POLICY = "fallback"
|
|
14
|
+
BLOCK_BYTES = 8 * 1024
|
|
15
|
+
TURKISH = frozenset("çğıöşüÇĞİÖŞÜ")
|
|
16
|
+
TYPOGRAPHY = str.maketrans({"’": "'", "‘": "'", "ʼ": "'", "´": "'", "`": "'", "“": '"', "”": '"', "„": '"',
|
|
17
|
+
"«": '"', "»": '"', "–": "-", "—": "-", "−": "-", "…": "..."})
|
|
18
|
+
UNSAFE = re.compile("[\x00-\x08\x0b-\x1f\x7f-\x9f\u061c\u200e\u200f\u202a-\u202e\u2066-\u2069]")
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class Frontend:
|
|
22
|
+
def __init__(self, vocab):
|
|
23
|
+
self.vocab = frozenset(vocab)
|
|
24
|
+
self.normalizer = Normalizer()
|
|
25
|
+
|
|
26
|
+
def __call__(self, text):
|
|
27
|
+
text = UNSAFE.sub(" ", text.encode("utf-8", "ignore").decode("utf-8"))
|
|
28
|
+
if not text.strip():
|
|
29
|
+
return ""
|
|
30
|
+
return self.alphabet(" ".join(self.spoken(block) for block in blocks(text)))
|
|
31
|
+
|
|
32
|
+
def spoken(self, text):
|
|
33
|
+
try:
|
|
34
|
+
return self.normalizer.normalize(text, ambiguity_policy=POLICY).normalized_text
|
|
35
|
+
except Exception: # invalid input or a resource limit: keep the words rather than lose the sentence
|
|
36
|
+
return text
|
|
37
|
+
|
|
38
|
+
def alphabet(self, text):
|
|
39
|
+
text = text.translate(TYPOGRAPHY).replace("İ", "i").replace("I", "ı").lower()
|
|
40
|
+
out = []
|
|
41
|
+
for ch in text:
|
|
42
|
+
if ch not in TURKISH:
|
|
43
|
+
ch = "".join(c for c in unicodedata.normalize("NFKD", ch) if not unicodedata.combining(c))
|
|
44
|
+
out.append(ch if ch and all(c in self.vocab for c in ch) else " ")
|
|
45
|
+
return re.sub(r"\s+", " ", "".join(out)).strip()
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def blocks(text):
|
|
49
|
+
"""Split at whitespace into pieces the normalizer accepts in one call."""
|
|
50
|
+
words, block, size = text.split(), [], 0
|
|
51
|
+
for word in words:
|
|
52
|
+
n = len(word.encode("utf-8")) + 1
|
|
53
|
+
if block and size + n > BLOCK_BYTES:
|
|
54
|
+
yield " ".join(block)
|
|
55
|
+
block, size = [], 0
|
|
56
|
+
block.append(word)
|
|
57
|
+
size += n
|
|
58
|
+
if block:
|
|
59
|
+
yield " ".join(block)
|