LGTM / lgtm /modules.py
qnguyen3's picture
LGTM: PyTorch + ONNX weights and inference code
409d4fb verified
Raw History Blame Contribute Delete
13.4 kB
"""LGTM building blocks.
Tensor layout: channels-first (B, C, T); masks are (B, 1, T) float tensors of 0/1.
"""
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
# --------------------------------------------------------------------------
# Basic layers
# --------------------------------------------------------------------------
class ChannelLayerNorm(nn.Module):
"""LayerNorm over channels of a (B, C, T) tensor. ONNX eps = 1e-6."""
def __init__(self, dim, eps=1e-6):
super().__init__()
self.norm = nn.LayerNorm(dim, eps=eps)
def forward(self, x):
return self.norm(x.transpose(1, 2)).transpose(1, 2)
class Linear(nn.Module):
"""nn.Linear wrapped as ``.linear`` to match original names (``W_query.linear.weight``)."""
def __init__(self, idim, odim, bias=True):
super().__init__()
self.linear = nn.Linear(idim, odim, bias=bias)
def forward(self, x):
return self.linear(x)
class PaddedConv1d(nn.Module):
"""Conv1d with replicate ('edge') padding, wrapped as ``.net`` like the original.
causal=True pads (k-1)*d on the left only (autoencoder decoder),
otherwise (k-1)*d/2 on both sides (text encoder / vector field).
"""
def __init__(self, idim, odim, ksz, dilation=1, groups=1, bias=True, causal=False):
super().__init__()
self.net = nn.Conv1d(idim, odim, ksz, dilation=dilation, groups=groups, bias=bias)
total = (ksz - 1) * dilation
self.pad = (total, 0) if causal else (total // 2, total - total // 2)
def forward(self, x):
if self.pad != (0, 0):
x = F.pad(x, self.pad, mode="replicate")
return self.net(x)
class ConvNeXtBlock(nn.Module):
"""ConvNeXt-1D block: dwconv -> LN -> pw(4x) -> GELU(erf) -> pw -> gamma, residual.
With a mask: input, dwconv output and block output are multiplied by the
mask (exactly as in the text encoder / vector-field graphs). The decoder
uses no mask and causal padding.
"""
def __init__(self, dim, intermediate_dim, ksz, dilation=1, causal=False, wrapped_dwconv=False):
super().__init__()
self.gamma = nn.Parameter(torch.full((1, dim, 1), 1e-6))
dw = PaddedConv1d(dim, dim, ksz, dilation=dilation, groups=dim, causal=causal)
if wrapped_dwconv:
self.dwconv = dw # params: dwconv.net.{weight,bias}
else:
# params: dwconv.{weight,bias}; keep padding logic in the block.
self.dwconv = dw.net
self._pad = dw.pad
self.wrapped = wrapped_dwconv
self.norm = ChannelLayerNorm(dim)
self.pwconv1 = nn.Conv1d(dim, intermediate_dim, 1)
self.act = nn.GELU() # exact erf GELU, as in the graph
self.pwconv2 = nn.Conv1d(intermediate_dim, dim, 1)
def forward(self, x, mask=None):
if mask is not None:
x = x * mask
residual = x
if self.wrapped:
y = self.dwconv(x)
else:
y = self.dwconv(F.pad(x, self._pad, mode="replicate"))
if mask is not None:
y = y * mask
y = self.norm(y)
y = self.pwconv2(self.act(self.pwconv1(y)))
x = residual + self.gamma * y
if mask is not None:
x = x * mask
return x
class ConvNeXtStack(nn.Module):
"""Stack of ConvNeXt blocks, params at ``convnext.{i}.*``."""
def __init__(self, idim, ksz, intermediate_dim, num_layers, dilation_lst, causal=False, wrapped_dwconv=False, **_):
super().__init__()
assert len(dilation_lst) == num_layers
self.convnext = nn.ModuleList(
[
ConvNeXtBlock(idim, intermediate_dim, ksz, d, causal=causal, wrapped_dwconv=wrapped_dwconv)
for d in dilation_lst
]
)
def forward(self, x, mask=None):
for blk in self.convnext:
x = blk(x, mask)
return x
# --------------------------------------------------------------------------
# VITS-style relative-position self-attention encoder (text / sentence enc.)
# --------------------------------------------------------------------------
class RelPosMultiHeadAttention(nn.Module):
"""VITS MultiHeadAttention with shared relative position embeddings (window 4)."""
def __init__(self, channels, n_heads, window_size=4):
super().__init__()
assert channels % n_heads == 0
self.n_heads = n_heads
self.k_channels = channels // n_heads
self.window_size = window_size
self.conv_q = nn.Conv1d(channels, channels, 1)
self.conv_k = nn.Conv1d(channels, channels, 1)
self.conv_v = nn.Conv1d(channels, channels, 1)
self.conv_o = nn.Conv1d(channels, channels, 1)
std = self.k_channels ** -0.5
self.emb_rel_k = nn.Parameter(torch.randn(1, 2 * window_size + 1, self.k_channels) * std)
self.emb_rel_v = nn.Parameter(torch.randn(1, 2 * window_size + 1, self.k_channels) * std)
def forward(self, x, attn_mask):
q, k, v = self.conv_q(x), self.conv_k(x), self.conv_v(x)
b, d, t = k.shape
h, kc = self.n_heads, self.k_channels
query = q.view(b, h, kc, t).transpose(2, 3) / math.sqrt(kc)
key = k.view(b, h, kc, t).transpose(2, 3)
value = v.view(b, h, kc, t).transpose(2, 3)
scores = torch.matmul(query, key.transpose(-2, -1))
key_rel = self._get_relative_embeddings(self.emb_rel_k, t)
rel_logits = torch.matmul(query, key_rel.unsqueeze(0).transpose(-2, -1))
scores = scores + self._relative_to_absolute(rel_logits)
scores = scores.masked_fill(attn_mask == 0, -1e4)
p = F.softmax(scores, dim=-1)
out = torch.matmul(p, value)
value_rel = self._get_relative_embeddings(self.emb_rel_v, t)
out = out + torch.matmul(self._absolute_to_relative(p), value_rel.unsqueeze(0))
out = out.transpose(2, 3).contiguous().view(b, d, t)
return self.conv_o(out)
def _get_relative_embeddings(self, emb, length):
w = self.window_size
pad_length = max(length - (w + 1), 0)
start = max((w + 1) - length, 0)
emb = F.pad(emb, (0, 0, pad_length, pad_length, 0, 0)) # pad of 0 is a no-op (keeps export branch-free)
return emb[:, start: start + 2 * length - 1]
@staticmethod
def _relative_to_absolute(x):
b, h, l, _ = x.shape
x = F.pad(x, (0, 1))
x = x.view(b, h, l * 2 * l)
x = F.pad(x, (0, l - 1))
return x.view(b, h, l + 1, 2 * l - 1)[:, :, :l, l - 1:]
@staticmethod
def _absolute_to_relative(x):
b, h, l, _ = x.shape
x = F.pad(x, (0, l - 1))
x = x.view(b, h, l * l + l * (l - 1))
x = F.pad(x, (l, 0))
return x.view(b, h, l, 2 * l)[:, :, :, 1:]
class FFN(nn.Module):
def __init__(self, channels, filter_channels):
super().__init__()
self.conv_1 = nn.Conv1d(channels, filter_channels, 1)
self.conv_2 = nn.Conv1d(filter_channels, channels, 1)
def forward(self, x, mask):
x = torch.relu(self.conv_1(x * mask))
return self.conv_2(x * mask) * mask
class AttnEncoder(nn.Module):
"""VITS post-norm transformer encoder."""
def __init__(self, hidden_channels, filter_channels, n_heads, n_layers, p_dropout=0.0):
super().__init__()
self.attn_layers = nn.ModuleList([RelPosMultiHeadAttention(hidden_channels, n_heads) for _ in range(n_layers)])
self.norm_layers_1 = nn.ModuleList([ChannelLayerNorm(hidden_channels) for _ in range(n_layers)])
self.ffn_layers = nn.ModuleList([FFN(hidden_channels, filter_channels) for _ in range(n_layers)])
self.norm_layers_2 = nn.ModuleList([ChannelLayerNorm(hidden_channels) for _ in range(n_layers)])
self.drop = nn.Dropout(p_dropout)
def forward(self, x, mask):
attn_mask = mask.unsqueeze(2) * mask.unsqueeze(-1)
x = x * mask
for attn, n1, ffn, n2 in zip(self.attn_layers, self.norm_layers_1, self.ffn_layers, self.norm_layers_2):
x = n1(x + self.drop(attn(x, attn_mask)))
x = n2(x + self.drop(ffn(x, mask)))
return x * mask
class CharEmbedder(nn.Module):
def __init__(self, n_vocab, dim):
super().__init__()
self.char_embedder = nn.Embedding(n_vocab, dim)
def forward(self, ids, mask):
return self.char_embedder(ids).transpose(1, 2) * mask
# --------------------------------------------------------------------------
# Style (GST-like) cross attention: keys go through tanh.
# --------------------------------------------------------------------------
class StyleAttention(nn.Module):
"""Multi-head cross-attention to style tokens (used in the text encoder and
the vector field's style-conditioning layers).
Heads are formed by splitting the last dim and stacking on a new leading
axis; keys are passed through tanh; scores are divided by ``sqrt(n_units)``;
rows of padded queries are zeroed after softmax.
"""
def __init__(self, q_dim, k_dim, v_dim, n_units, n_heads, out_dim):
super().__init__()
self.n_heads = n_heads
self.n_units = n_units
self.W_query = Linear(q_dim, n_units)
self.W_key = Linear(k_dim, n_units)
self.W_value = Linear(v_dim, n_units)
self.out_fc = Linear(n_units, out_dim)
def _heads(self, x): # (B, T, U) -> (H, B, T, U/H)
return torch.stack(torch.chunk(x, self.n_heads, dim=-1), dim=0)
def forward(self, q, k, v, q_mask=None):
"""q: (B, Tq, q_dim), k: (B, Tk, k_dim), v: (B, Tk, v_dim), q_mask: (B, Tq, 1)."""
q = self._heads(self.W_query(q))
k = torch.tanh(self._heads(self.W_key(k)))
v = self._heads(self.W_value(v))
p = F.softmax(torch.matmul(q, k.transpose(-1, -2)) / math.sqrt(self.n_units), dim=-1)
if q_mask is not None:
p = p * q_mask.unsqueeze(0)
out = torch.matmul(p, v) # (H, B, Tq, d)
out = torch.cat(out.unbind(0), dim=-1)
return self.out_fc(out)
# --------------------------------------------------------------------------
# Rotary cross attention (latent queries -> text keys), vector field.
# --------------------------------------------------------------------------
class RotaryCrossAttention(nn.Module):
"""Text-conditioning attention in the vector field.
Rotary embedding uses *length-normalised* positions: pos = i / length,
angle = pos * theta, theta_j = rotary_scale * base^(-j/(d/2)). Queries use
the latent length, keys the text length (so the attention learns a soft
monotonic alignment in relative position). Non-interleaved rotation
(first half / second half). Scores are divided by sqrt(n_units / 2).
"""
def __init__(self, idim, text_dim, n_units, n_heads, rotary_base=10000, rotary_scale=10, **_):
super().__init__()
self.n_heads = n_heads
self.head_dim = n_units // n_heads
self.scale = math.sqrt(n_units / 2) # = 16 for n_units=512 (verified against ONNX)
self.W_query = Linear(idim, n_units)
self.W_key = Linear(text_dim, n_units)
self.W_value = Linear(text_dim, n_units)
self.out_fc = Linear(n_units, idim)
half = self.head_dim // 2
theta = rotary_scale * rotary_base ** (-torch.arange(half, dtype=torch.float32) / half)
self.register_buffer("theta", theta.view(1, 1, half))
def _heads(self, x): # (B, T, U) -> (H, B, T, d)
b, t, _ = x.shape
return x.view(b, t, self.n_heads, self.head_dim).permute(2, 0, 1, 3)
def _rotate(self, x, mask):
# x: (H, B, T, d); mask: (B, T, 1)
t = x.shape[2]
length = mask.sum(dim=(1, 2)).view(-1, 1, 1)
pos = torch.arange(t, device=x.device, dtype=x.dtype).view(1, t, 1) / length
ang = pos * self.theta # (B, T, d/2)
cos, sin = torch.cos(ang), torch.sin(ang)
x1, x2 = x.chunk(2, dim=-1)
return torch.cat([x1 * cos - x2 * sin, x1 * sin + x2 * cos], dim=-1)
def forward(self, x, text, x_mask, text_mask):
"""x: (B, Tq, idim), text: (B, Tk, text_dim), masks (B, T, 1)."""
q = self._rotate(self._heads(self.W_query(x)), x_mask)
k = self._rotate(self._heads(self.W_key(text)), text_mask)
v = self._heads(self.W_value(text))
scores = torch.matmul(q, k.transpose(-1, -2)) / self.scale
key_mask = text_mask.transpose(1, 2).unsqueeze(0) # (1, B, 1, Tk)
scores = scores.masked_fill(key_mask == 0, float("-inf"))
p = F.softmax(scores, dim=-1) * x_mask.unsqueeze(0)
out = torch.matmul(p, v) # (H, B, Tq, d)
b, tq = out.shape[1], out.shape[2]
out = out.permute(1, 2, 0, 3).reshape(b, tq, -1)
return self.out_fc(out)
class TimeEncoder(nn.Module):
"""sinusoidal(t * 1000) -> Linear -> Mish -> Linear."""
def __init__(self, time_dim, hdim):
super().__init__()
half = time_dim // 2
freqs = torch.exp(-math.log(10000) * torch.arange(half, dtype=torch.float32) / (half - 1))
self.register_buffer("freqs", freqs.view(1, half))
self.mlp = nn.Sequential(Linear(time_dim, hdim), nn.Mish(), Linear(hdim, time_dim))
def forward(self, t): # t: (B,) in [0, 1]
ang = t.view(-1, 1) * 1000.0 * self.freqs
return self.mlp(torch.cat([torch.sin(ang), torch.cos(ang)], dim=-1))