"""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))