Download Modules/diffusion/pitch_dit.py from FashionFlora/SFlowTTS: direct link, hf CLI and curl.
- Browser
- Download file 17.4 kB
-
https://huggingface.co/FashionFlora/SFlowTTS/resolve/main/Modules/diffusion/pitch_dit.py
- Command line
-
hf download hf://FashionFlora/SFlowTTS/Modules/diffusion/pitch_dit.py
-
curl -L -o pitch_dit.py https://huggingface.co/FashionFlora/SFlowTTS/resolve/main/Modules/diffusion/pitch_dit.py
17.4 kB
| 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 | |
| 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) |