SFlowTTS / Modules /diffusion /style_transformer.py
FashionFlora's picture
Upload full repo excluding dump_40, dump_100, precomputed_tokens, precomputed_data
fb0011a verified
Raw History Blame Contribute Delete
18.6 kB
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Tuple
from einops import repeat
from einops.layers.torch import Rearrange
# -----------------------------------------------------------------------------
# small utils (keep yours if already present)
# -----------------------------------------------------------------------------
def exists(val):
return val is not None
def rand_bool(shape, proba, device=None):
if proba == 1:
return torch.ones(shape, device=device, dtype=torch.bool)
if proba == 0:
return torch.zeros(shape, device=device, dtype=torch.bool)
return torch.bernoulli(torch.full(shape, proba, device=device)).to(torch.bool)
# -----------------------------------------------------------------------------
# time embedding (reuse your existing ones if already in this file)
# -----------------------------------------------------------------------------
class LearnedPositionalEmbedding(nn.Module):
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: torch.Tensor) -> torch.Tensor:
# x: [B]
x = x.view(-1, 1)
freqs = x * self.weights.view(1, -1) * 2 * math.pi
fouriered = torch.cat((freqs.sin(), freqs.cos()), dim=-1)
return torch.cat((x, fouriered), dim=-1)
def TimePositionalEmbedding(dim: int, out_features: int) -> nn.Module:
return nn.Sequential(
LearnedPositionalEmbedding(dim),
nn.Linear(dim + 1, out_features),
)
class FixedEmbedding(nn.Module):
"""
Safer than your current version: supports any length by wrapping positions.
"""
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: torch.Tensor) -> torch.Tensor:
b, t = x.shape[0], x.shape[1]
device = x.device
pos = torch.arange(t, device=device) % self.max_length
emb = self.embedding(pos) # [T, D]
return repeat(emb, "t d -> b t d", b=b)
# -----------------------------------------------------------------------------
# HRM building blocks (RoPE attention + SwiGLU FFN + AdaRMSNorm)
# -----------------------------------------------------------------------------
class AdaRMSNorm(nn.Module):
def __init__(self, style_dim: int, channels: int, eps: float = 1e-5):
super().__init__()
self.fc = nn.Linear(style_dim, channels * 2)
self.eps = eps
with torch.no_grad():
self.fc.weight.zero_()
self.fc.bias.zero_()
def forward(self, x: torch.Tensor, style: torch.Tensor) -> torch.Tensor:
# x: [B, T, C], style: [B, S]
gb = self.fc(style)
gamma, beta = gb.chunk(2, dim=-1)
gamma = gamma.unsqueeze(1).to(dtype=x.dtype)
beta = beta.unsqueeze(1).to(dtype=x.dtype)
x_norm = x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
return (1.0 + gamma) * x_norm + beta
class SwiGLUFFN(nn.Module):
def __init__(self, dim: int, expansion: float = 2.0, dropout: float = 0.0):
super().__init__()
hidden = int(dim * expansion)
self.fc = nn.Linear(dim, hidden * 2)
self.proj = nn.Linear(hidden, dim)
self.dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
a, b = self.fc(x).chunk(2, dim=-1)
x = F.silu(b) * a
x = self.proj(self.dropout(x))
return x
class RotaryEmbedding(nn.Module):
def __init__(self, dim: int, base: float = 10000.0):
super().__init__()
if dim % 2 != 0:
raise ValueError("RoPE head_dim must be even.")
inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))
self.register_buffer("inv_freq", inv_freq, persistent=False)
def get_cos_sin(
self,
seq_len: int,
device: torch.device,
dtype: torch.dtype,
) -> Tuple[torch.Tensor, torch.Tensor]:
t = torch.arange(seq_len, device=device, dtype=dtype)
freqs = torch.outer(t, self.inv_freq.to(dtype))
emb = torch.cat([freqs, freqs], dim=-1) # [T, head_dim]
cos = emb.cos().unsqueeze(0).unsqueeze(0) # [1, 1, T, D]
sin = emb.sin().unsqueeze(0).unsqueeze(0)
return cos, sin
def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
# x: [B, H, T, D]
x_even = x[..., ::2]
x_odd = x[..., 1::2]
x_rot = torch.stack([-x_odd, x_even], dim=-1).flatten(-2)
return (x * cos) + (x_rot * sin)
class MultiheadSelfAttentionRoPE(nn.Module):
def __init__(
self,
embed_dim: int,
num_heads: int,
dropout: float = 0.0,
use_rope: bool = True,
rope_base: float = 10000.0,
):
super().__init__()
if embed_dim % num_heads != 0:
raise ValueError("embed_dim must be divisible by num_heads.")
self.embed_dim = embed_dim
self.num_heads = num_heads
self.head_dim = embed_dim // num_heads
self.scale = 1.0 / math.sqrt(self.head_dim)
self.qkv = nn.Linear(embed_dim, embed_dim * 3)
self.out_proj = nn.Linear(embed_dim, embed_dim)
self.attn_dropout = nn.Dropout(dropout)
self.use_rope = use_rope
self.rope = RotaryEmbedding(self.head_dim, base=rope_base) if use_rope else None
if use_rope and (self.head_dim % 2 != 0):
raise ValueError(
f"RoPE requires even head_dim, got head_dim={self.head_dim}."
)
def forward(
self,
x: torch.Tensor,
key_padding_mask: Optional[torch.Tensor] = None,
) -> torch.Tensor:
# x: [B, T, C], key_padding_mask: [B, T] with True = PAD
b, t, c = x.shape
q, k, v = self.qkv(x).chunk(3, dim=-1)
q = q.view(b, t, self.num_heads, self.head_dim).transpose(1, 2)
k = k.view(b, t, self.num_heads, self.head_dim).transpose(1, 2)
v = v.view(b, t, self.num_heads, self.head_dim).transpose(1, 2)
if self.use_rope:
cos, sin = self.rope.get_cos_sin(t, x.device, x.dtype)
q = apply_rope(q, cos, sin)
k = apply_rope(k, cos, sin)
attn_scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale
if key_padding_mask is not None:
pad = key_padding_mask.unsqueeze(1).unsqueeze(2) # [B,1,1,T]
attn_scores = attn_scores.masked_fill(
pad, torch.finfo(attn_scores.dtype).min
)
attn = F.softmax(attn_scores, dim=-1)
attn = self.attn_dropout(attn)
out = torch.matmul(attn, v) # [B,H,T,D]
out = out.transpose(1, 2).contiguous().view(b, t, c)
return self.out_proj(out)
class HRMReasoningBlock(nn.Module):
def __init__(
self,
dim: int,
num_heads: int,
expansion: float,
style_dim: int,
attn_dropout: float = 0.0,
resid_dropout: float = 0.0,
use_rope: bool = True,
rope_base: float = 10000.0,
rms_eps: float = 1e-5,
):
super().__init__()
self.attn = MultiheadSelfAttentionRoPE(
dim, num_heads, attn_dropout, use_rope, rope_base
)
self.ff = SwiGLUFFN(dim, expansion, resid_dropout)
self.resid_dropout = nn.Dropout(resid_dropout)
self.ada1 = AdaRMSNorm(style_dim, dim, eps=rms_eps)
self.ada2 = AdaRMSNorm(style_dim, dim, eps=rms_eps)
def forward(
self,
hidden_states: torch.Tensor,
style: torch.Tensor,
key_padding_mask: Optional[torch.Tensor] = None,
) -> torch.Tensor:
attn_out = self.attn(hidden_states, key_padding_mask=key_padding_mask)
hidden_states = self.ada1(
hidden_states + self.resid_dropout(attn_out), style
)
mlp_out = self.ff(hidden_states)
hidden_states = self.ada2(
hidden_states + self.resid_dropout(mlp_out), style
)
return hidden_states
class HRMReasoningModule(nn.Module):
def __init__(self, layers: nn.ModuleList):
super().__init__()
self.layers = layers
def forward(
self,
hidden_states: torch.Tensor,
input_injection: torch.Tensor,
style: torch.Tensor,
key_padding_mask: Optional[torch.Tensor] = None,
) -> torch.Tensor:
hidden_states = hidden_states + input_injection
for layer in self.layers:
hidden_states = layer(
hidden_states, style, key_padding_mask=key_padding_mask
)
return hidden_states
# -----------------------------------------------------------------------------
# HRM-based StyleTransformer1d (same external API as your current one)
# -----------------------------------------------------------------------------
class StyleTransformer1d(nn.Module):
"""
HRM replacement for your StyleTransformer1d.
- Works on tokens: concat([x, embedding]) => model_dim = channels + cond_dim
- Uses HRM H/L cycles instead of stacked transformer blocks
- Style-conditioning via AdaRMSNorm in every block
- Time/features mapping added as a bias/injection
"""
def __init__(
self,
num_layers: int,
channels: int,
num_heads: int,
head_features: int, # kept for signature compatibility; ignored
multiplier: int,
use_context_time: bool = True,
use_rel_pos: bool = False, # ignored
context_features_multiplier: int = 1, # ignored
rel_pos_num_buckets: Optional[int] = None, # ignored
rel_pos_max_distance: Optional[int] = None, # ignored
context_features: Optional[int] = None, # style dim
context_embedding_features: Optional[int] = None, # cond dim
embedding_max_length: int = 512,
H_cycles: int = 1,
L_cycles: int = 2,
attn_dropout: float = 0.0,
resid_dropout: float = 0.0,
use_rope: bool = True,
rope_base: float = 10000.0,
):
super().__init__()
assert exists(context_embedding_features), "context_embedding_features is required"
assert exists(context_features), "context_features (style dim) is required"
self.channels = channels
self.cond_dim = context_embedding_features
self.style_dim = context_features
self.model_dim = channels + context_embedding_features
self.use_context_time = use_context_time
self.use_context_features = exists(context_features)
self.H_cycles = H_cycles
self.L_cycles = L_cycles
# Optional positional embedding added to conditioning sequence
self.fixed_embedding = FixedEmbedding(
max_length=embedding_max_length,
features=context_embedding_features,
)
# time/features mapping -> model_dim (same as concat(x, embedding))
if use_context_time or self.use_context_features:
self.to_mapping = nn.Sequential(
nn.Linear(self.model_dim, self.model_dim),
nn.GELU(),
nn.Linear(self.model_dim, self.model_dim),
nn.GELU(),
)
if use_context_time:
self.to_time = nn.Sequential(
TimePositionalEmbedding(dim=channels, out_features=self.model_dim),
nn.GELU(),
)
if self.use_context_features:
self.to_features = nn.Sequential(
nn.Linear(context_features, self.model_dim),
nn.GELU(),
)
# HRM states + step embeddings
self.H_init = nn.Parameter(torch.zeros(1, 1, self.model_dim))
self.L_init = nn.Parameter(torch.zeros(1, 1, self.model_dim))
nn.init.trunc_normal_(self.H_init, std=0.02)
nn.init.trunc_normal_(self.L_init, std=0.02)
self.h_step_emb = nn.Embedding(H_cycles, self.model_dim)
self.l_step_emb = nn.Embedding(L_cycles, self.model_dim)
layers_L = nn.ModuleList(
[
HRMReasoningBlock(
dim=self.model_dim,
num_heads=num_heads,
expansion=float(multiplier),
style_dim=self.style_dim,
attn_dropout=attn_dropout,
resid_dropout=resid_dropout,
use_rope=use_rope,
rope_base=rope_base,
)
for _ in range(num_layers)
]
)
layers_H = nn.ModuleList(
[
HRMReasoningBlock(
dim=self.model_dim,
num_heads=num_heads,
expansion=float(multiplier),
style_dim=self.style_dim,
attn_dropout=attn_dropout,
resid_dropout=resid_dropout,
use_rope=use_rope,
rope_base=rope_base,
)
for _ in range(num_layers)
]
)
self.L_level = HRMReasoningModule(layers_L)
self.H_level = HRMReasoningModule(layers_H)
# project back to `channels`
self.to_out = nn.Sequential(
Rearrange("b t c -> b c t"),
nn.Conv1d(self.model_dim, channels, kernel_size=1),
)
def get_mapping(
self,
time: Optional[torch.Tensor] = None,
features: Optional[torch.Tensor] = None,
) -> Optional[torch.Tensor]:
items = []
if self.use_context_time:
assert exists(time), "use_context_time=True but time was not provided"
items.append(self.to_time(time))
if self.use_context_features:
assert exists(features), "context_features exists but features were not provided"
items.append(self.to_features(features))
if len(items) == 0:
return None
mapping = torch.stack(items, dim=0).sum(dim=0) # [B, model_dim]
return self.to_mapping(mapping)
def run(
self,
x: torch.Tensor, # [B, T, channels]
time: torch.Tensor, # [B]a
embedding: torch.Tensor, # [B, T, cond_dim]
features: torch.Tensor, # [B, style_dim]
key_padding_mask: Optional[torch.Tensor] = None, # [B, T] True=PAD
) -> torch.Tensor:
# add positional embedding to cond (this was computed but unused in your old code)
embedding = embedding + self.fixed_embedding(embedding).to(dtype=embedding.dtype)
# concat token-wise
x_cat = torch.cat([x, embedding], dim=-1) # [B, T, model_dim]
mapping = self.get_mapping(time=time, features=features)
if mapping is not None:
x_cat = x_cat + mapping.unsqueeze(1)
b, t, _ = x_cat.shape
z_H = self.H_init.to(dtype=x_cat.dtype).expand(b, t, -1)
z_L = self.L_init.to(dtype=x_cat.dtype).expand(b, t, -1)
for h in range(self.H_cycles):
h_emb = self.h_step_emb.weight[h].view(1, 1, -1).to(dtype=x_cat.dtype)
z_H_curr = z_H + h_emb
for l in range(self.L_cycles):
l_emb = self.l_step_emb.weight[l].view(1, 1, -1).to(dtype=x_cat.dtype)
inj = z_H_curr + x_cat + l_emb
z_L = self.L_level(
z_L,
input_injection=inj,
style=features,
key_padding_mask=key_padding_mask,
)
if key_padding_mask is not None:
z_L = z_L.masked_fill(key_padding_mask.unsqueeze(-1), 0.0)
z_H = self.H_level(
z_H,
input_injection=z_L,
style=features,
key_padding_mask=key_padding_mask,
)
if key_padding_mask is not None:
z_H = z_H.masked_fill(key_padding_mask.unsqueeze(-1), 0.0)
out = self.to_out(z_H).transpose(1, 2) # [B, T, channels]
return out
def forward(
self,
x: torch.Tensor,
time: torch.Tensor,
embedding_mask_proba: float = 0.0,
embedding: Optional[torch.Tensor] = None,
features: Optional[torch.Tensor] = None,
embedding_scale: float = 1.0,
key_padding_mask: Optional[torch.Tensor] = None,
) -> torch.Tensor:
assert exists(embedding), "embedding is required"
assert exists(features), "features(style) is required"
b = embedding.shape[0]
device = embedding.device
if embedding_mask_proba > 0.0:
# Shape [B, 1, 1] for broadcasting
batch_mask = rand_bool((b, 1, 1), embedding_mask_proba, device=device)
# Mask the embedding [B, T, D]
embedding = torch.where(batch_mask, torch.zeros_like(embedding), embedding)
# Mask the features [B, S] -> squeeze mask to [B, 1] for broadcasting
features_mask = batch_mask.squeeze(-1) # [B, 1]
features = torch.where(features_mask, torch.zeros_like(features), features)
if embedding_scale != 1.0:
out = self.run(
x=x,
time=time,
embedding=embedding,
features=features,
key_padding_mask=key_padding_mask,
)
null_embedding = torch.zeros_like(embedding)
out_null = self.run(
x=x,
time=time,
embedding=null_embedding,
features=features,
key_padding_mask=key_padding_mask,
)
return out_null + (out - out_null) * embedding_scale
return self.run(
x=x,
time=time,
embedding=embedding,
features=features,
key_padding_mask=key_padding_mask,
)