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