import torch import torch.nn as nn import torch.nn.functional as F from einops import rearrange from torch import Tensor from math import log, pi from typing import Callable, Optional, Tuple # ----------------------------------------------------------------------------- # Helper functions # ----------------------------------------------------------------------------- def exists(x) -> bool: return x is not None def default(val, fn): return val if val is not None else fn() def rand_bool(shape, proba: float, device: torch.device) -> Tensor: """Sample Bernoulli mask with given probability of True.""" return torch.rand(shape, device=device) < proba # ----------------------------------------------------------------------------- # Normalization / Transformer components # ----------------------------------------------------------------------------- class RMSNorm(nn.Module): """RMSNorm (does not subtract mean).""" def __init__( self, dim: int, eps: float = 1e-8, elementwise_affine: bool = True, ): super().__init__() self.dim = dim self.eps = eps self.elementwise_affine = elementwise_affine if elementwise_affine: self.weight = nn.Parameter(torch.ones(dim)) else: self.register_buffer("weight", torch.ones(dim)) def forward(self, x: Tensor) -> Tensor: # x: (..., dim) rms = x.pow(2).mean(dim=-1, keepdim=True).add(self.eps).sqrt() x_normed = x / rms return x_normed * self.weight def _make_norm(norm_type: str, dim: int) -> nn.Module: """Create a norm module by name.""" if norm_type is None or norm_type == "none": return nn.Identity() if norm_type == "layer": return nn.LayerNorm(dim) if norm_type == "rms": return RMSNorm(dim) raise ValueError(f"Unknown norm_type: {norm_type}") class RotaryEmbedding(nn.Module): """ RoPE implementation that generates sin/cos on demand. Works on head dimension d (must be even). """ def __init__(self, dim: int, max_seq_len: int = 512, base: int = 10000): super().__init__() assert dim % 2 == 0, "Rotary embedding dim must be even" self.dim = dim self.max_seq_len = max_seq_len self.base = base inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim)) self.register_buffer("inv_freq", inv_freq, persistent=False) def _build_sin_cos( self, seq_len: int, device: torch.device, dtype: torch.dtype, ) -> Tuple[Tensor, Tensor]: positions = torch.arange(seq_len, device=device, dtype=dtype).unsqueeze(1) angles = positions * rearrange( self.inv_freq.to(device=device, dtype=dtype), "d -> 1 d", ) sin = torch.sin(angles) cos = torch.cos(angles) sin = torch.stack([sin, sin], dim=-1).reshape(seq_len, self.dim) cos = torch.stack([cos, cos], dim=-1).reshape(seq_len, self.dim) return sin, cos @staticmethod def rotate_half(x: Tensor) -> Tensor: x1 = x[..., ::2] x2 = x[..., 1::2] x_rotated = torch.stack((-x2, x1), dim=-1).reshape_as(x) return x_rotated def apply_rotary(self, q: Tensor, k: Tensor) -> Tuple[Tensor, Tensor]: """ q, k: (b, h, n, d) with d == self.dim """ assert q.shape[-1] == self.dim and k.shape[-1] == self.dim seq_len = q.shape[-2] device = q.device dtype = q.dtype sin, cos = self._build_sin_cos(seq_len, device=device, dtype=dtype) sin = sin.unsqueeze(0).unsqueeze(0) # [1,1,n,d] cos = cos.unsqueeze(0).unsqueeze(0) q_out = (q * cos) + (self.rotate_half(q) * sin) k_out = (k * cos) + (self.rotate_half(k) * sin) return q_out, k_out def FeedForward(features: int, multiplier: int) -> nn.Module: mid_features = features * multiplier return nn.Sequential( nn.Linear(features, mid_features), nn.GELU(), nn.Linear(mid_features, features), ) class AttentionBase(nn.Module): def __init__( self, features: int, *, head_features: int, num_heads: int, use_rope: bool, rope_max_seq_len: int = 512, out_features: Optional[int] = None, attn_dropout: float = 0.0, out_dropout: float = 0.0, ): super().__init__() self.scale = head_features ** -0.5 self.num_heads = num_heads self.use_rope = use_rope mid_features = head_features * num_heads if out_features is None: out_features = features self.to_out = nn.Linear(mid_features, out_features) self.attn_dropout = ( nn.Dropout(attn_dropout) if attn_dropout > 0.0 else nn.Identity() ) self.out_dropout = ( nn.Dropout(out_dropout) if out_dropout > 0.0 else nn.Identity() ) if use_rope: self.rotary = RotaryEmbedding( dim=head_features, max_seq_len=rope_max_seq_len, ) else: self.rotary = None def forward(self, q: Tensor, k: Tensor, v: Tensor, mask: Optional[Tensor] = None) -> Tensor: # q,k,v: (b, n, h*d) q = rearrange(q, "b n (h d) -> b h n d", h=self.num_heads) k = rearrange(k, "b n (h d) -> b h n d", h=self.num_heads) v = rearrange(v, "b n (h d) -> b h n d", h=self.num_heads) if self.rotary is not None: q, k = self.rotary.apply_rotary(q, k) # sim: [Batch, Heads, QueryLen, KeyLen] sim = torch.einsum("... n d, ... m d -> ... n m", q, k) sim = sim * self.scale # --- MASKING APPLIED HERE --- if mask is not None: # mask should be broadcastable to [Batch, Heads, QueryLen, KeyLen] # Usually mask is [B, 1, 1, KeyLen] where True=Keep, False=Mask # We use a large negative number instead of -inf to be safe with fp16 mixed precision mask_value = -torch.finfo(sim.dtype).max if sim.dtype == torch.float32 else -1e4 sim = sim.masked_fill(~mask, mask_value) # ---------------------------- attn = sim.softmax(dim=-1) attn = self.attn_dropout(attn) out = torch.einsum("... n m, ... m d -> ... n d", attn, v) out = rearrange(out, "b h n d -> b n (h d)") out = self.to_out(out) out = self.out_dropout(out) return out class Attention(nn.Module): def __init__( self, features: int, *, head_features: int, num_heads: int, out_features: Optional[int] = None, use_rope: bool, rope_max_seq_len: int = 512, attn_dropout: float = 0.0, out_dropout: float = 0.0, norm_type: str = "layer", ): super().__init__() mid_features = head_features * num_heads self.norm = _make_norm(norm_type, features) self.to_q = nn.Linear(features, mid_features, bias=False) self.to_kv = nn.Linear(features, mid_features * 2, bias=False) self.attention = AttentionBase( features=features, out_features=out_features, num_heads=num_heads, head_features=head_features, use_rope=use_rope, rope_max_seq_len=rope_max_seq_len, attn_dropout=attn_dropout, out_dropout=out_dropout, ) def forward(self, x: Tensor, mask: Optional[Tensor] = None) -> Tensor: if not isinstance(self.norm, nn.Identity): x_norm = self.norm(x) else: x_norm = x q = self.to_q(x_norm) kv = self.to_kv(x_norm) k, v = torch.chunk(kv, 2, dim=-1) # Pass mask down return self.attention(q, k, v, mask=mask) class TransformerBlock(nn.Module): def __init__( self, features: int, num_heads: int, head_features: int, multiplier: int, use_rope: bool, rope_max_seq_len: int = 512, dropout: float = 0.0, attn_dropout: float = 0.0, ff_dropout: float = 0.0, norm_type: str = "layer", ): super().__init__() self.attention = Attention( features=features, num_heads=num_heads, head_features=head_features, use_rope=use_rope, rope_max_seq_len=rope_max_seq_len, attn_dropout=attn_dropout, out_dropout=dropout, norm_type=norm_type, ) self.feed_forward = FeedForward(features=features, multiplier=multiplier) self.norm_ff = _make_norm(norm_type, features) self.ff_dropout = ( nn.Dropout(ff_dropout) if ff_dropout > 0.0 else nn.Identity() ) def forward(self, x: Tensor, mask: Optional[Tensor] = None) -> Tensor: # Pass mask to attention x = self.attention(x, mask=mask) + x if not isinstance(self.norm_ff, nn.Identity): ff_in = self.norm_ff(x) else: ff_in = x ff_out = self.feed_forward(ff_in) ff_out = self.ff_dropout(ff_out) x = ff_out + x return x # ----------------------------------------------------------------------------- # Time and fixed positional embeddings # ----------------------------------------------------------------------------- class LearnedPositionalEmbedding(nn.Module): """Continuous time embedding.""" def __init__(self, dim: int): super().__init__() assert dim % 2 == 0 half_dim = dim // 2 self.weights = nn.Parameter(torch.randn(half_dim)) def forward(self, x: Tensor) -> Tensor: x = rearrange(x, "b -> b 1") freqs = x * rearrange(self.weights, "d -> 1 d") * 2 * pi fouriered = torch.cat((freqs.sin(), freqs.cos()), dim=-1) fouriered = torch.cat((x, fouriered), dim=-1) return fouriered def TimePositionalEmbedding(dim: int, out_features: int) -> nn.Module: return nn.Sequential( LearnedPositionalEmbedding(dim), nn.Linear(in_features=dim + 1, out_features=out_features), ) class FixedEmbedding(nn.Module): def __init__(self, max_length: int, features: int): super().__init__() self.max_length = max_length self.embedding = nn.Embedding(max_length, features) def forward(self, x: Tensor) -> Tensor: """ x: [B, T, D] (only T and B are used, D is ignored) returns: [B, T, features] """ batch_size, length = x.shape[0], x.shape[1] device = x.device assert ( length <= self.max_length ), "Input sequence length must be <= max_length" position = torch.arange(length, device=device) fixed_embedding = self.embedding(position) # [T, features] fixed_embedding = fixed_embedding.unsqueeze(0).expand( batch_size, length, -1, ) return fixed_embedding # ----------------------------------------------------------------------------- # Seq-to-seq diffusion net for joint pitch + energy # ----------------------------------------------------------------------------- class PitchEnergyTransformer1d(nn.Module): """ Seq-to-seq transformer used as diffusion network. Input: x : [B, 2, T] (2 channels: pitch, energy) time : [B] (noise level embedding, from VKDiffusion) embedding : [B, 512, T] (per-frame conditioning, e.g. P from predictor) mask : [B, T] (Boolean padding mask, True=Keep, False=Pad) Output: v_pred : [B, 2, T] """ def __init__( self, *, num_layers: int = 6, num_heads: int = 8, head_features: int = 64, multiplier: int = 4, ctx_dim: int = 512, max_seq_len: int = 2048, use_rope: bool = True, dropout: float = 0.0, attn_dropout: float = 0.0, ff_dropout: float = 0.0, norm_type: str = "rms", ): super().__init__() self.data_channels = 2 self.ctx_dim = ctx_dim self.max_seq_len = max_seq_len total_features = self.data_channels + ctx_dim self.input_dropout = ( nn.Dropout(dropout) if dropout > 0.0 else nn.Identity() ) self.blocks = nn.ModuleList( [ TransformerBlock( features=total_features, num_heads=num_heads, head_features=head_features, multiplier=multiplier, use_rope=use_rope, rope_max_seq_len=max_seq_len, dropout=dropout, attn_dropout=attn_dropout, ff_dropout=ff_dropout, norm_type=norm_type, ) for _ in range(num_layers) ] ) self.to_out = nn.Linear(total_features, self.data_channels) # time embedding self.to_time = nn.Sequential( TimePositionalEmbedding( dim=self.data_channels, out_features=total_features, ), nn.GELU(), ) self.to_mapping = nn.Sequential( nn.Linear(total_features, total_features), nn.GELU(), nn.Linear(total_features, total_features), nn.GELU(), ) self.fixed_embedding = FixedEmbedding( max_length=max_seq_len, features=ctx_dim, ) def get_mapping(self, time: Tensor) -> Tensor: mapping = self.to_time(time) mapping = self.to_mapping(mapping) return mapping # [B, F] def _run_core( self, x: Tensor, time: Tensor, ctx: Tensor, mask: Optional[Tensor] = None ) -> Tensor: """ x : [B, 2, T] ctx : [B, ctx_dim, T] mask: [B, T] (True=Keep, False=Pad) """ b, c, t = x.shape assert c == self.data_channels assert ( ctx.shape[0] == b and ctx.shape[1] == self.ctx_dim and ctx.shape[2] == t ) x_seq = x.transpose(1, 2) # [B, T, 2] ctx_seq = ctx.transpose(1, 2) # [B, T, ctx_dim] h = torch.cat([x_seq, ctx_seq], dim=-1) # [B, T, F] mapping = self.get_mapping(time).unsqueeze(1) # [B,1,F] h = self.input_dropout(h) # --- PREPARE ATTENTION MASK --- # If mask is present, we reshape it for broadcasting in AttentionBase # Input mask: [B, T] # Target shape for attention: [B, 1, 1, T] (Broadcasting over Heads and QueryLen) attn_mask = None if mask is not None: attn_mask = mask.unsqueeze(1).unsqueeze(1) # [B, 1, 1, T] # ------------------------------ for block in self.blocks: h = block(h + mapping, mask=attn_mask) h = self.to_out(h) # [B, T, 2] out = h.transpose(1, 2) # [B, 2, T] return out def forward( self, x: Tensor, time: Tensor, *, mask: Optional[Tensor] = None, embedding: Tensor, embedding_mask_proba: float = 0.0, embedding_scale: float = 1.0, ) -> Tensor: """ x : [B, 2, T] time : [B] embedding : [B, 512, T] mask : [B, T] (Padding mask. True=Valid Frame, False=Pad) """ assert embedding is not None b, ctx_dim, t = embedding.shape assert ctx_dim == self.ctx_dim device = embedding.device emb_seq = embedding.transpose(1, 2) # [B, T, ctx_dim] fixed_emb_seq = self.fixed_embedding(emb_seq) # [B, T, ctx_dim] # Apply dropout to the *conditioning embedding* (Classifier-Free Guidance style) if embedding_mask_proba > 0.0: rand_mask = rand_bool( shape=(b, 1, 1), proba=embedding_mask_proba, device=device, ) emb_seq = torch.where(rand_mask, fixed_emb_seq, emb_seq) emb_ctx = emb_seq.transpose(1, 2) # [B, ctx_dim, T] fixed_ctx = fixed_emb_seq.transpose(1, 2) # [B, ctx_dim, T] if embedding_scale != 1.0: # CFG: We must pass the padding mask to both calls out_cond = self._run_core(x, time, ctx=emb_ctx, mask=mask) out_uncond = self._run_core(x, time, ctx=fixed_ctx, mask=mask) return out_uncond + (out_cond - out_uncond) * embedding_scale return self._run_core(x, time, ctx=emb_ctx, mask=mask)