SFlowTTS / Modules /diffusion /pitch_dit.py
FashionFlora's picture
Upload full repo excluding dump_40, dump_100, precomputed_tokens, precomputed_data
fb0011a verified
Raw History Blame Contribute Delete
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
@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)