"""LGTM text-to-speech model (44.1 kHz, 11 languages, zero-shot voice cloning). LGTM ttl text-to-latent: text encoder, style encoder, flow-matching vector field ae speech autoencoder: encoder (audio -> latent) and decoder (latent -> 44.1 kHz audio) dp utterance-level duration predictor (+ its style encoder) Latents: raw AE latent (B, 24, T_f) at 44100/512 Hz is normalised and 6 frames are stacked into channels -> (B, 144, T_f / 6); the flow-matching model works in that space. """ import torch import torch.nn as nn import torch.nn.functional as F from .modules import ( AttnEncoder, ChannelLayerNorm, CharEmbedder, ConvNeXtStack, Linear, PaddedConv1d, RotaryCrossAttention, StyleAttention, TimeEncoder, ) N_VOCAB = 8322 # ========================================================================== # Text-to-latent # ========================================================================== class TextEncoder(nn.Module): """char emb -> ConvNeXt x6 -> relpos transformer x4, with a long residual.""" def __init__(self, cfg): super().__init__() self.text_embedder = CharEmbedder(N_VOCAB, cfg["text_embedder"]["char_emb_dim"]) self.convnext = ConvNeXtStack(**cfg["convnext"]) self.attn_encoder = AttnEncoder(**cfg["attn_encoder"]) # cfg["proj_out"] is an identity projection. def forward(self, text_ids, text_mask): x = self.text_embedder(text_ids, text_mask) x = self.convnext(x, text_mask) x = self.attn_encoder(x, text_mask) + x return x * text_mask class SpeechPromptedTextEncoder(nn.Module): """Two cross-attentions from text to the 50 style tokens; both residuals are taken w.r.t. the *original* text features (as in the graph).""" def __init__(self, text_dim, style_dim, n_units, n_heads): super().__init__() self.attention1 = StyleAttention(text_dim, style_dim, style_dim, n_units, n_heads, text_dim) self.attention2 = StyleAttention(text_dim, style_dim, style_dim, n_units, n_heads, text_dim) self.norm = ChannelLayerNorm(text_dim) def forward(self, x, text_mask, style_key, style_value): xt = x.transpose(1, 2) m = text_mask.transpose(1, 2) h = self.attention1(xt, style_key, style_value, m) * m + xt h = self.attention2(h, style_key, style_value, m) * m + xt return self.norm.norm(h).transpose(1, 2) * text_mask class StyleTokenLayer(nn.Module): """Reference encoder pooling -> n_style tokens. """ def __init__(self, input_dim, n_style, style_key_dim, style_value_dim, prototype_dim, n_units, n_heads): super().__init__() if style_key_dim > 0: self.style_key = nn.Parameter(torch.randn(1, n_style, style_key_dim) * 0.02) else: self.style_key = None self.prototype = nn.Parameter(torch.randn(1, n_style, prototype_dim) * 0.02) self.attention = StyleAttention(prototype_dim, input_dim, input_dim, n_units, n_heads, style_value_dim) # Output = normalize(base + delta): `base` is the mean voice style, the network predicts the # voice-specific deviation. self.base = nn.Parameter(torch.zeros(1, n_style, style_value_dim)) nn.init.zeros_(self.attention.out_fc.linear.weight) nn.init.zeros_(self.attention.out_fc.linear.bias) def forward(self, h, mask): """h: (B, C, T) encoded reference, mask: (B, 1, T) -> (B, n_style, style_value_dim).""" b = h.shape[0] ht = h.transpose(1, 2) q = self.prototype.expand(b, -1, -1) # mask padded reference frames by pushing their keys to a neutral value # and removing their values; StyleAttention has no key mask of its own. return self._masked_attention(q, ht, mask.transpose(1, 2)) def _masked_attention(self, q, kv, kv_mask): att = self.attention qh = att._heads(att.W_query(q)) kh = torch.tanh(att._heads(att.W_key(kv))) vh = att._heads(att.W_value(kv)) scores = torch.matmul(qh, kh.transpose(-1, -2)) / (att.n_units ** 0.5) scores = scores.masked_fill(kv_mask.transpose(1, 2).unsqueeze(0) == 0, float("-inf")) out = torch.matmul(F.softmax(scores, dim=-1), vh) out = self.base + att.out_fc(torch.cat(out.unbind(0), dim=-1)) # voice styles (ttl 50x256 and dp 8x16) are per-token unit vectors return F.normalize(out, dim=-1) class StyleEncoder(nn.Module): """Reference latent (B,144,T) -> style tokens.""" def __init__(self, cfg): super().__init__() p = cfg["proj_in"] self.proj_in = PaddedConv1d(p["ldim"] * p["chunk_compress_factor"], p["odim"], 1) self.convnext = ConvNeXtStack(**cfg["convnext"]) self.style_token_layer = StyleTokenLayer(**cfg["style_token_layer"]) def forward(self, latent, mask): h = self.proj_in(latent) * mask h = self.convnext(h, mask) return self.style_token_layer(h, mask) class UncondMasker(nn.Module): """Learned null tokens for classifier-free guidance.""" def __init__(self, text_dim, n_style, style_key_dim, style_value_dim, **_): super().__init__() self.text_special_token = nn.Parameter(torch.zeros(1, text_dim, 1)) self.style_key_special_token = nn.Parameter(torch.zeros(1, n_style, style_key_dim)) self.style_value_special_token = nn.Parameter(torch.zeros(1, n_style, style_value_dim)) class TimeCondBlock(nn.Module): def __init__(self, idim, time_dim): super().__init__() self.linear = Linear(time_dim, idim) def forward(self, x, mask, t_emb): return (x + self.linear(t_emb).unsqueeze(-1)) * mask class TextCondBlock(nn.Module): def __init__(self, idim, text_dim, n_heads, n_units, rotary_base, rotary_scale, **_): super().__init__() self.attn = RotaryCrossAttention(idim, text_dim, n_units, n_heads, rotary_base, rotary_scale) self.norm = ChannelLayerNorm(idim) def forward(self, x, mask, text_emb, text_mask): x = x * mask y = self.attn(x.transpose(1, 2), text_emb.transpose(1, 2), mask.transpose(1, 2), text_mask.transpose(1, 2)) x = x + y.transpose(1, 2) * mask return self.norm(x) * mask class StyleCondBlock(nn.Module): def __init__(self, idim, style_dim, n_units=256, n_heads=2): super().__init__() self.attention = StyleAttention(idim, style_dim, style_dim, n_units, n_heads, idim) self.norm = ChannelLayerNorm(idim) def forward(self, x, mask, style_key, style_value): x = x * mask m = mask.transpose(1, 2) y = self.attention(x.transpose(1, 2), style_key, style_value, m) * m x = x + y.transpose(1, 2) return self.norm(x) * mask class VectorField(nn.Module): """Flow-matching velocity estimator on compressed latents (B, 144, T).""" def __init__(self, cfg): super().__init__() p = cfg["proj_in"] ldim = p["ldim"] * p["chunk_compress_factor"] self.proj_in = PaddedConv1d(ldim, p["odim"], 1, bias=False) self.time_encoder = TimeEncoder(cfg["time_encoder"]["time_dim"], cfg["time_encoder"]["hdim"]) mb = cfg["main_blocks"] blocks = [] for _ in range(mb["n_blocks"]): blocks += [ ConvNeXtStack(**mb["convnext_0"]), TimeCondBlock(**mb["time_cond_layer"]), ConvNeXtStack(**mb["convnext_1"]), TextCondBlock(**mb["text_cond_layer"]), ConvNeXtStack(**mb["convnext_2"]), StyleCondBlock(**mb["style_cond_layer"]), ] self.main_blocks = nn.ModuleList(blocks) self.last_convnext = ConvNeXtStack(**cfg["last_convnext"]) self.proj_out = PaddedConv1d(p["odim"], ldim, 1, bias=False) def forward(self, x, t, text_emb, text_mask, style_key, style_value, latent_mask): t_emb = self.time_encoder(t) h = self.proj_in(x) * latent_mask for blk in self.main_blocks: if isinstance(blk, ConvNeXtStack): h = blk(h, latent_mask) elif isinstance(blk, TimeCondBlock): h = blk(h, latent_mask, t_emb) elif isinstance(blk, TextCondBlock): h = blk(h, latent_mask, text_emb, text_mask) else: h = blk(h, latent_mask, style_key, style_value) h = self.last_convnext(h, latent_mask) return self.proj_out(h) * latent_mask class TextToLatent(nn.Module): def __init__(self, cfg): super().__init__() self.cfg = cfg self.text_encoder = TextEncoder(cfg["text_encoder"]) self.style_encoder = StyleEncoder(cfg["style_encoder"]) s = cfg["speech_prompted_text_encoder"] self.speech_prompted_text_encoder = SpeechPromptedTextEncoder(s["text_dim"], s["style_dim"], s["n_units"], s["n_heads"]) self.uncond_masker = UncondMasker(**cfg["uncond_masker"]) self.vector_field = VectorField(cfg["vector_field"]) self.sig_min = cfg["flow_matching"]["sig_min"] @property def style_key(self): return self.style_encoder.style_token_layer.style_key def encode_text(self, text_ids, text_mask, style_ttl): """== text_encoder.onnx. Returns text_emb (B, 256, T).""" key = self.style_key.expand(text_ids.shape[0], -1, -1) x = self.text_encoder(text_ids, text_mask) return self.speech_prompted_text_encoder(x, text_mask, key, style_ttl) def velocity(self, x, t, text_emb, text_mask, style_ttl, latent_mask, uncond=False): """Single (conditional or unconditional) velocity prediction.""" b = x.shape[0] if uncond: um = self.uncond_masker text_emb = um.text_special_token.expand(b, -1, text_emb.shape[-1]) key = um.style_key_special_token.expand(b, -1, -1) style_ttl = um.style_value_special_token.expand(b, -1, -1) else: key = self.style_key.expand(b, -1, -1) return self.vector_field(x, t, text_emb, text_mask, key, style_ttl, latent_mask) def cfg_velocity(self, x, t, text_emb, text_mask, style_ttl, latent_mask, cfg_scale=4.0): """CFG as baked into vector_estimator.onnx: v_u + 4 (v_c - v_u) = 4 v_c - 3 v_u.""" xx = torch.cat([x, x]) tt = torch.cat([t, t]) te = torch.cat([text_emb, self.uncond_masker.text_special_token.expand_as(text_emb)]) tm = torch.cat([text_mask, text_mask]) lm = torch.cat([latent_mask, latent_mask]) b = x.shape[0] key = torch.cat([self.style_key.expand(b, -1, -1), self.uncond_masker.style_key_special_token.expand(b, -1, -1)]) val = torch.cat([style_ttl, self.uncond_masker.style_value_special_token.expand(b, -1, -1)]) v = self.vector_field(xx, tt, te, tm, key, val, lm) v_c, v_u = v[:b], v[b:] # (slicing instead of chunk keeps the batch dim dynamic in ONNX) return v_u + cfg_scale * (v_c - v_u) def euler_step(self, x, step, total_step, text_emb, text_mask, style_ttl, latent_mask, cfg_scale=4.0): """== vector_estimator.onnx (one Euler step on t = step / total_step).""" t = step / total_step v = self.cfg_velocity(x, t, text_emb, text_mask, style_ttl, latent_mask, cfg_scale) return (x + v / total_step.view(-1, 1, 1)) * latent_mask # ========================================================================== # Speech autoencoder # ========================================================================== class LatentDecoder(nn.Module): """== vocoder.onnx decoder: causal ConvNeXt, head outputs 512 samples/frame.""" def __init__(self, cfg): super().__init__() h = cfg["hdim"] self.embed = PaddedConv1d(cfg["idim"], h, cfg["ksz_init"], causal=True) self.convnext = ConvNeXtStack( h, cfg["ksz"], cfg["intermediate_dim"], cfg["num_layers"], cfg["dilation_lst"], causal=True, wrapped_dwconv=True ).convnext self.final_norm = nn.Module() self.final_norm.norm = nn.BatchNorm1d(h, eps=1e-5) hd = cfg["head"] self.head = nn.Module() self.head.layer1 = PaddedConv1d(hd["idim"], hd["hdim"], hd["ksz"], causal=True) self.head.act = nn.PReLU(1) self.head.layer2 = nn.Conv1d(hd["hdim"], hd["odim"], 1, bias=False) def forward(self, z): """z: raw latent (B, 24, T_f) -> wav (B, T_f * 512).""" x = self.embed(z) for blk in self.convnext: x = blk(x) x = self.final_norm.norm(x) x = self.head.layer2(self.head.act(self.head.layer1(x))) # (B, 512, T_f) return x.transpose(1, 2).reshape(x.shape[0], -1) class SpecProcessor(nn.Module): """log |STFT| (1025) ++ log mel (228) = 1253 features at hop 512. Frames are left-aligned to the decoder: frame i covers samples ending at 512*(i+1) (left pad n_fft - hop), so frame count == ceil(len / 512). """ def __init__(self, n_fft, win_length, hop_length, n_mels, sample_rate, eps, **_): super().__init__() import torchaudio self.n_fft, self.hop, self.win = n_fft, hop_length, win_length self.eps = eps self.register_buffer("window", torch.hann_window(win_length), persistent=False) fb = torchaudio.functional.melscale_fbanks(n_fft // 2 + 1, 0.0, sample_rate / 2, n_mels, sample_rate, norm="slaney", mel_scale="slaney") self.register_buffer("mel_fb", fb, persistent=False) def forward(self, wav): n = wav.shape[-1] n_frames = (n + self.hop - 1) // self.hop wav = F.pad(wav, (self.n_fft - self.hop, n_frames * self.hop - n)) spec = torch.stft(wav, self.n_fft, self.hop, self.win, self.window, center=False, return_complex=True).abs() mel = torch.matmul(spec.transpose(1, 2), self.mel_fb).transpose(1, 2) return torch.cat([torch.log(spec + self.eps), torch.log(mel + self.eps)], dim=1) class LatentEncoder(nn.Module): """wav @ 44.1 kHz -> raw latent (B, 24, T_f).""" def __init__(self, cfg): super().__init__() self.spec_processor = SpecProcessor(**cfg["spec_processor"]) h = cfg["hdim"] self.embed = PaddedConv1d(cfg["idim"], h, cfg["ksz_init"]) self.convnext = ConvNeXtStack(h, cfg["ksz"], cfg["intermediate_dim"], cfg["num_layers"], cfg["dilation_lst"], wrapped_dwconv=True).convnext self.final_norm = ChannelLayerNorm(h) self.proj_out = nn.Conv1d(h, cfg["odim"], 1) def forward(self, wav): x = self.embed(self.spec_processor(wav)) for blk in self.convnext: x = blk(x) return self.proj_out(self.final_norm(x)) class SpeechAutoencoder(nn.Module): def __init__(self, cfg, ttl_cfg): super().__init__() self.sample_rate = cfg["sample_rate"] self.hop = cfg["base_chunk_size"] self.ccf = ttl_cfg["chunk_compress_factor"] self.ldim = cfg["ldim"] self.register_buffer("latent_mean", torch.zeros(1, cfg["ldim"], 1)) self.register_buffer("latent_std", torch.ones(1, cfg["ldim"], 1)) self.register_buffer("normalizer_scale", torch.tensor(float(ttl_cfg["normalizer"]["scale"]))) self.encoder = LatentEncoder(cfg["encoder"]) self.decoder = LatentDecoder(cfg["decoder"]) # ---- latent <-> TTL representation ----------------------------------- def compress(self, z): """raw (B, 24, T_f) -> TTL latent (B, 144, ceil(T_f/6)); pads with the latent mean.""" b, c, t = z.shape zn = (z - self.latent_mean) / self.latent_std pad = (-t) % self.ccf if pad: zn = F.pad(zn, (0, pad)) zn = zn.view(b, c, -1, self.ccf).permute(0, 1, 3, 2).reshape(b, c * self.ccf, -1) return zn * self.normalizer_scale def decompress(self, x): """TTL latent (B, 144, T) -> raw (B, 24, 6T).""" b, _, t = x.shape x = x / self.normalizer_scale z = x.view(b, self.ldim, self.ccf, t).permute(0, 1, 3, 2).reshape(b, self.ldim, t * self.ccf) return z * self.latent_std + self.latent_mean def decode_ttl(self, x): """== vocoder.onnx: TTL latent (B, 144, T) -> wav (B, 3072 T).""" return self.decoder(self.decompress(x)) def encode_ttl(self, wav): """wav (B, N) @ 44.1 kHz -> TTL latent (B, 144, ceil(N/3072)).""" return self.compress(self.encoder(wav)) # ========================================================================== # Duration predictor # ========================================================================== class SentenceEncoder(nn.Module): def __init__(self, cfg): super().__init__() d = cfg["char_emb_dim"] self.sentence_token = nn.Parameter(torch.randn(1, d, 1) * 0.02) self.text_embedder = CharEmbedder(N_VOCAB, cfg["text_embedder"]["char_emb_dim"]) self.convnext = ConvNeXtStack(**cfg["convnext"]) self.attn_encoder = AttnEncoder(**cfg["attn_encoder"]) self.proj_out = PaddedConv1d(cfg["proj_out"]["idim"], cfg["proj_out"]["odim"], 1, bias=False) def forward(self, text_ids, text_mask): b = text_ids.shape[0] x = self.text_embedder(text_ids, text_mask) x = torch.cat([self.sentence_token.expand(b, -1, -1), x], dim=-1) m = torch.cat([torch.ones_like(text_mask[:, :, :1]), text_mask], dim=-1) x = self.convnext(x, m) x = self.attn_encoder(x, m) + x return (self.proj_out(x[:, :, :1]) * m[:, :, :1]).flatten(1) # (B, 64) class DurationHead(nn.Module): def __init__(self, sentence_dim, n_style, style_dim, hdim, n_layer): super().__init__() assert n_layer == 2 self.layers = nn.ModuleList([nn.Linear(sentence_dim + n_style * style_dim, hdim), nn.Linear(hdim, 1)]) self.activation = nn.PReLU(1) def forward(self, s, style_dp): h = torch.cat([s, style_dp.flatten(1)], dim=1) return self.layers[1](self.activation(self.layers[0](h))).squeeze(1) # log-seconds class DurationPredictor(nn.Module): def __init__(self, cfg): super().__init__() self.sentence_encoder = SentenceEncoder(cfg["sentence_encoder"]) self.style_encoder = StyleEncoder(cfg["style_encoder"]) self.predictor = DurationHead(**cfg["predictor"]) def log_duration(self, text_ids, text_mask, style_dp): return self.predictor(self.sentence_encoder(text_ids, text_mask), style_dp) def forward(self, text_ids, text_mask, style_dp): """== duration_predictor.onnx: total utterance duration in seconds (B,).""" return torch.exp(self.log_duration(text_ids, text_mask, style_dp)) # ========================================================================== class LGTM(nn.Module): def __init__(self, cfg): super().__init__() self.cfg = cfg self.ttl = TextToLatent(cfg["ttl"]) self.ae = SpeechAutoencoder(cfg["ae"], cfg["ttl"]) self.dp = DurationPredictor(cfg["dp"])