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