Download lgtm/model.py from polyskill/LGTM: direct link, hf CLI and curl.
- Browser
- Download file 19.1 kB
-
https://huggingface.co/polyskill/LGTM/resolve/main/lgtm/model.py
- Command line
-
hf download hf://polyskill/LGTM/lgtm/model.py
-
curl -L -o model.py https://huggingface.co/polyskill/LGTM/resolve/main/lgtm/model.py
19.1 kB
| """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"] | |
| 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"]) | |