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.
@@ -0,0 +1,5 @@
1
+ """EMA Lightning: tiny, fast and accurate Turkish text to speech."""
2
+ from .api import EMA, Speech
3
+
4
+ __all__ = ["EMA", "Speech"]
5
+ __version__ = "1.0.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())
@@ -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(",;:- ") + "."
@@ -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)
@@ -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)