LGTM / lgtm /model.py
qnguyen3's picture
LGTM: PyTorch + ONNX weights and inference code
409d4fb verified
Raw History Blame Contribute Delete
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"]
@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"])