Download Modules/vae.py from FashionFlora/SFlowTTS: direct link, hf CLI and curl.
- Browser
- Download file 16.3 kB
-
https://huggingface.co/FashionFlora/SFlowTTS/resolve/main/Modules/vae.py
- Command line
-
hf download hf://FashionFlora/SFlowTTS/Modules/vae.py
-
curl -L -o vae.py https://huggingface.co/FashionFlora/SFlowTTS/resolve/main/Modules/vae.py
16.3 kB
| # cvae_gan_latent.py | |
| # ------------------------------------------------------------ | |
| # Text + speaker + language conditioned latent CVAE + GAN to | |
| # predict 3 style latents (acoustic, pitch, prosodic) from text. | |
| # | |
| # - Teacher: StyleEncoderVAE_OLD (unchanged, external) | |
| # - Condition: text encoder tokens + speaker_id + language_id | |
| # - Generator: 3-layer Transformer over text tokens -> 3 styles | |
| # - Discriminator: multi-MLP, least-squares GAN + feat. match | |
| # | |
| # After training, you can discard the style encoder and use | |
| # the generator to produce style latents directly from | |
| # text + speaker_id + language_id. | |
| # ------------------------------------------------------------ | |
| from typing import Dict, List, Optional, Tuple | |
| import math | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from transformers import AutoModel, AutoTokenizer | |
| from Modules.diffusion.modules import * | |
| # ---------------- GAN losses (as provided) ----------------- | |
| def feature_loss(fmap_r, fmap_g): | |
| loss = 0 | |
| for dr, dg in zip(fmap_r, fmap_g): | |
| for rl, gl in zip(dr, dg): | |
| loss += torch.mean(torch.abs(rl - gl)) | |
| return loss * 2 | |
| def discriminator_loss(disc_real_outputs, disc_generated_outputs): | |
| loss = 0 | |
| r_losses = [] | |
| g_losses = [] | |
| for dr, dg in zip(disc_real_outputs, disc_generated_outputs): | |
| r_loss = torch.mean((1 - dr) ** 2) | |
| g_loss = torch.mean(dg**2) | |
| loss += r_loss + g_loss | |
| r_losses.append(r_loss.item()) | |
| g_losses.append(g_loss.item()) | |
| return loss, r_losses, g_losses | |
| def generator_loss(disc_outputs): | |
| loss = 0 | |
| gen_losses = [] | |
| for dg in disc_outputs: | |
| l = torch.mean((1 - dg) ** 2) | |
| gen_losses.append(l) | |
| loss += l | |
| return loss, gen_losses | |
| def masked_mean_pool( | |
| x: torch.Tensor, mask: Optional[torch.Tensor] | |
| ) -> torch.Tensor: | |
| """ | |
| x: [B, T, C] | |
| mask: [B, T] with 1 for valid, 0 for pad. If None, mean over T. | |
| returns: [B, C] | |
| """ | |
| if mask is None: | |
| return x.mean(dim=1) | |
| mask = mask.float() | |
| denom = torch.clamp(mask.sum(dim=1, keepdim=True), min=1.0) | |
| return (x * mask.unsqueeze(-1)).sum(dim=1) / denom | |
| class SinusoidalPositionalEncoding(nn.Module): | |
| def __init__(self, d_model: int, max_len: int = 512): | |
| super().__init__() | |
| pe = torch.zeros(max_len, d_model) | |
| pos = torch.arange(0, max_len, dtype=torch.float32).unsqueeze(1) | |
| div = torch.exp( | |
| torch.arange(0, d_model, 2, dtype=torch.float32) | |
| * (-math.log(10000.0) / d_model) | |
| ) | |
| pe[:, 0::2] = torch.sin(pos * div) | |
| pe[:, 1::2] = torch.cos(pos * div) | |
| self.register_buffer("pe", pe.unsqueeze(0), persistent=False) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| # x: [B, T, C] | |
| T = x.size(1) | |
| return x + self.pe[:, :T, :] | |
| # --------------- 3-style Transformer Generator ------------- | |
| class StyleLatentGenerator(nn.Module): | |
| def __init__( | |
| self, | |
| cond_dim: int = 128, | |
| style_dim: int = 128, | |
| n_styles: int = 3, | |
| n_layers: int = 3, | |
| n_heads: int = 4, | |
| head_features: int = 32, | |
| ff_mult: int = 4, | |
| max_len: int = 512, | |
| dropout: float = 0.1, | |
| attn_dropout: float = 0.0, | |
| ff_dropout: float = 0.0, | |
| use_rope: bool = False, | |
| rope_max_seq_len: int = 512, | |
| norm_type: str = "layer", | |
| embedding_mask_proba: float = 0.0, | |
| # Speaker / Language Config | |
| num_languages: int = 0, | |
| max_speakers_per_language: int = 0, # Added this | |
| spk_emb_dim: Optional[int] = None, | |
| lang_emb_dim: Optional[int] = None, | |
| ): | |
| super().__init__() | |
| self.cond_dim = cond_dim | |
| self.style_dim = style_dim | |
| self.total_style_dim = style_dim * n_styles | |
| self.embedding_mask_proba = embedding_mask_proba | |
| # --- Logic Fix: Calculate Total Speakers --- | |
| self.num_languages = num_languages | |
| self.max_speakers_per_language = max_speakers_per_language | |
| # Calculate total unique embeddings needed | |
| if num_languages > 0 and max_speakers_per_language > 0: | |
| self.num_speakers_total = num_languages * max_speakers_per_language | |
| else: | |
| self.num_speakers_total = 0 | |
| if spk_emb_dim is None: spk_emb_dim = cond_dim | |
| if lang_emb_dim is None: lang_emb_dim = cond_dim | |
| self.spk_emb_dim = spk_emb_dim if self.num_speakers_total > 0 else 0 | |
| self.lang_emb_dim = lang_emb_dim if num_languages > 0 else 0 | |
| # Embeddings | |
| if self.num_speakers_total > 0: | |
| self.spk_embed = nn.Embedding(self.num_speakers_total, spk_emb_dim) | |
| else: | |
| self.spk_embed = None | |
| if num_languages > 0: | |
| self.lang_embed = nn.Embedding(num_languages, lang_emb_dim) | |
| else: | |
| self.lang_embed = None | |
| # Projection: [cond + spk + lang] -> cond_dim | |
| extra_cond_dim = self.spk_emb_dim + self.lang_emb_dim | |
| if extra_cond_dim > 0: | |
| self.cond_proj = nn.Linear(cond_dim + extra_cond_dim, cond_dim) | |
| else: | |
| self.cond_proj = None | |
| # Learned "Query" Token (CLS) | |
| self.cls = nn.Parameter(torch.randn(1, 1, self.total_style_dim) * 0.02) | |
| # Transformer | |
| self.transformer = Transformer1d( | |
| num_layers=n_layers, | |
| channels=self.total_style_dim, | |
| num_heads=n_heads, | |
| head_features=head_features, | |
| multiplier=ff_mult, | |
| use_context_time=False, | |
| use_rope=use_rope, | |
| rope_max_seq_len=rope_max_seq_len, | |
| context_embedding_features=cond_dim, # Cross-attention dim | |
| embedding_max_length=max_len, | |
| dropout=dropout, | |
| attn_dropout=attn_dropout, | |
| ff_dropout=ff_dropout, | |
| norm_type=norm_type, | |
| ) | |
| # Output Head | |
| self.to_style = nn.Sequential( | |
| nn.LayerNorm(self.total_style_dim), | |
| nn.Linear(self.total_style_dim, self.total_style_dim * 2), | |
| nn.GELU(), | |
| nn.Linear(self.total_style_dim * 2, self.total_style_dim), | |
| ) | |
| def _compute_global_speaker_ids(self, speaker_ids, language_ids): | |
| """Helper to map local speaker ID to global embedding index.""" | |
| if self.max_speakers_per_language <= 0: | |
| return None | |
| # Safety checks | |
| if speaker_ids.max() >= self.max_speakers_per_language: | |
| raise ValueError(f"Speaker ID exceeds max_speakers_per_language ({self.max_speakers_per_language})") | |
| return (language_ids * self.max_speakers_per_language) + speaker_ids | |
| def _fuse_condition( | |
| self, | |
| cond_tokens: torch.Tensor, | |
| cond_mask: Optional[torch.Tensor], | |
| speaker_ids: Optional[torch.Tensor], | |
| language_ids: Optional[torch.Tensor], | |
| ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: | |
| B, T, C = cond_tokens.shape | |
| tokens = cond_tokens | |
| if self.cond_proj is not None: | |
| extras = [] | |
| # Fuse Language | |
| if self.lang_embed is not None: | |
| assert language_ids is not None | |
| lang_vec = self.lang_embed(language_ids) # [B, D_l] | |
| lang = lang_vec.unsqueeze(1).expand(-1, T, -1) | |
| extras.append(lang) | |
| # Fuse Speaker (Global Offset) | |
| if self.spk_embed is not None: | |
| assert speaker_ids is not None and language_ids is not None | |
| # FIX: Calculate global ID | |
| spk_global = self._compute_global_speaker_ids(speaker_ids, language_ids) | |
| spk_vec = self.spk_embed(spk_global) # [B, D_s] | |
| spk = spk_vec.unsqueeze(1).expand(-1, T, -1) | |
| extras.append(spk) | |
| if extras: | |
| # Concatenate along channel dim: [Text | Lang | Spk] | |
| tokens = torch.cat([tokens] + extras, dim=-1) | |
| tokens = self.cond_proj(tokens) | |
| # Masking optimization | |
| if cond_mask is not None: | |
| valid_counts = cond_mask.sum(dim=1) | |
| max_valid = int(valid_counts.max().item()) | |
| if max_valid == 0: | |
| return torch.zeros_like(tokens), cond_mask | |
| tokens = tokens[:, :max_valid, :] | |
| mask = cond_mask[:, :max_valid] | |
| else: | |
| mask = None | |
| return tokens, mask | |
| def forward( | |
| self, | |
| cond_tokens: torch.Tensor, | |
| cond_mask: Optional[torch.Tensor] = None, | |
| speaker_ids: Optional[torch.Tensor] = None, | |
| language_ids: Optional[torch.Tensor] = None, | |
| ) -> torch.Tensor: | |
| B = cond_tokens.size(0) | |
| # 1. Prepare Text Condition (as Cross-Attention context) | |
| tokens, mask = self._fuse_condition( | |
| cond_tokens, cond_mask, speaker_ids, language_ids | |
| ) | |
| # 2. Prepare Latent Query (CLS token) | |
| cls = self.cls.expand(B, 1, -1) # [B, 1, total_style_dim] | |
| # 3. Transformer Logic | |
| # We pass 'tokens' as 'embedding'. | |
| # Crucial: This assumes Transformer1d performs Cross-Attention against 'embedding'. | |
| out = self.transformer.forward( | |
| cls, | |
| None, # time | |
| embedding_mask_proba=self.embedding_mask_proba, | |
| embedding=tokens, # Context | |
| embedding_scale=1.0, | |
| ) | |
| z_hat = out.squeeze(1) # [B, total_style_dim] | |
| z_hat = self.to_style(z_hat) | |
| return z_hat | |
| def split_3_styles(z_all: torch.Tensor, style_dim: int = 128): | |
| return torch.split(z_all, style_dim, dim=-1) | |
| class LatentDiscSub(nn.Module): | |
| """ | |
| Spectral-norm MLP that returns a logit and intermediate features. | |
| Input is [z || cond], where cond is a pooled condition vector. | |
| """ | |
| def __init__(self, in_dim: int, hidden_dims: List[int]): | |
| super().__init__() | |
| layers = [] | |
| last = in_dim | |
| for h in hidden_dims: | |
| linear = nn.utils.spectral_norm(nn.Linear(last, h)) | |
| layers += [linear] | |
| layers += [nn.LeakyReLU(0.2, inplace=True)] | |
| last = h | |
| self.mlp = nn.Sequential(*layers) | |
| self.final = nn.utils.spectral_norm(nn.Linear(last, 1)) | |
| def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, List[torch.Tensor]]: | |
| feats = [] | |
| cur = x | |
| for layer in self.mlp: | |
| cur = layer(cur) | |
| if isinstance(layer, nn.LeakyReLU): | |
| feats.append(cur) | |
| logit = self.final(cur) | |
| return logit, feats | |
| class MultiLatentDiscriminator(nn.Module): | |
| """ | |
| Wrapper around multiple sub-MLPs. | |
| Condition: | |
| - text tokens | |
| - language_id (for lang embedding) | |
| - speaker_id local to that language, turned into global speaker | |
| index the same way as in the generator. | |
| """ | |
| def __init__( | |
| self, | |
| z_dim: int, | |
| cond_dim: int, | |
| n_subs: int = 3, | |
| hidden_dims: Optional[List[int]] = None, | |
| cond_pool: str = "mean", | |
| dropout: float = 0.0, | |
| num_languages: int = 0, | |
| max_speakers_per_language: int = 0, | |
| spk_emb_dim: Optional[int] = None, | |
| lang_emb_dim: Optional[int] = None, | |
| ): | |
| super().__init__() | |
| if hidden_dims is None: | |
| hidden_dims = [128, 128, 64] | |
| self.cond_dim_tokens = cond_dim | |
| self.cond_pool = cond_pool | |
| self.dropout = nn.Dropout(p=dropout) | |
| self.num_languages = num_languages | |
| self.max_speakers_per_language = max_speakers_per_language | |
| if spk_emb_dim is None: | |
| spk_emb_dim = cond_dim | |
| if lang_emb_dim is None: | |
| lang_emb_dim = cond_dim | |
| # language embeddings | |
| if num_languages > 0: | |
| self.lang_embed = nn.Embedding(num_languages, lang_emb_dim) | |
| self.lang_emb_dim = lang_emb_dim | |
| else: | |
| self.lang_embed = None | |
| self.lang_emb_dim = 0 | |
| # speaker embeddings (language‑dependent) | |
| if num_languages > 0 and max_speakers_per_language > 0: | |
| num_speakers_total = num_languages * max_speakers_per_language | |
| self.spk_embed = nn.Embedding(num_speakers_total, spk_emb_dim) | |
| self.spk_emb_dim = spk_emb_dim | |
| self.num_speakers_total = num_speakers_total | |
| else: | |
| self.spk_embed = None | |
| self.spk_emb_dim = 0 | |
| self.num_speakers_total = 0 | |
| extra_dim = self.spk_emb_dim + self.lang_emb_dim | |
| if extra_dim > 0: | |
| self.cond_fuse = nn.Linear( | |
| self.cond_dim_tokens + extra_dim, self.cond_dim_tokens | |
| ) | |
| else: | |
| self.cond_fuse = None | |
| in_dim = z_dim + self.cond_dim_tokens | |
| self.subs = nn.ModuleList( | |
| [LatentDiscSub(in_dim=in_dim, hidden_dims=hidden_dims) for _ in range(n_subs)] | |
| ) | |
| def _compute_global_speaker_ids( | |
| self, | |
| speaker_ids: torch.Tensor, | |
| language_ids: torch.Tensor, | |
| ) -> torch.Tensor: | |
| assert ( | |
| self.max_speakers_per_language > 0 | |
| ), "max_speakers_per_language must be > 0 when using speakers." | |
| if speaker_ids.max().item() >= self.max_speakers_per_language: | |
| raise ValueError( | |
| f"speaker_ids contain value >= max_speakers_per_language " | |
| f"({self.max_speakers_per_language})." | |
| ) | |
| spk_global = ( | |
| language_ids * self.max_speakers_per_language + speaker_ids | |
| ) | |
| if spk_global.max().item() >= self.num_speakers_total: | |
| raise ValueError( | |
| "Computed global speaker id out of range in discriminator. " | |
| "Check num_languages and max_speakers_per_language." | |
| ) | |
| return spk_global | |
| def pool_cond( | |
| self, | |
| cond_tokens: torch.Tensor, | |
| cond_mask: Optional[torch.Tensor], | |
| speaker_ids: Optional[torch.Tensor], | |
| language_ids: Optional[torch.Tensor], | |
| ) -> torch.Tensor: | |
| cond_vec = masked_mean_pool(cond_tokens, cond_mask) # [B, C] | |
| extras = [] | |
| if self.lang_embed is not None: | |
| assert language_ids is not None, ( | |
| "language_ids must be provided when num_languages > 0." | |
| ) | |
| lang_vec = self.lang_embed(language_ids) # [B, D_l] | |
| extras.append(lang_vec) | |
| if self.spk_embed is not None: | |
| assert ( | |
| speaker_ids is not None and language_ids is not None | |
| ), "speaker_ids and language_ids must be given for speakers." | |
| spk_global = self._compute_global_speaker_ids( | |
| speaker_ids=speaker_ids, language_ids=language_ids | |
| ) # [B] | |
| spk_vec = self.spk_embed(spk_global) # [B, D_s] | |
| extras.append(spk_vec) | |
| if self.cond_fuse is not None and extras: | |
| cond_full = torch.cat([cond_vec] + extras, dim=-1) | |
| cond_vec = self.cond_fuse(cond_full) | |
| return cond_vec | |
| def forward( | |
| self, | |
| z: torch.Tensor, | |
| cond_tokens: torch.Tensor, | |
| cond_mask: Optional[torch.Tensor] = None, | |
| speaker_ids: Optional[torch.Tensor] = None, | |
| language_ids: Optional[torch.Tensor] = None, | |
| ) -> Tuple[List[torch.Tensor], List[List[torch.Tensor]]]: | |
| cond_vec = self.pool_cond( | |
| cond_tokens, cond_mask, speaker_ids, language_ids | |
| ) # [B, C] | |
| x = torch.cat([z, cond_vec], dim=-1) | |
| x = self.dropout(x) | |
| logits = [] | |
| features = [] | |
| for sub in self.subs: | |
| logit, feats = sub(x) | |
| logits.append(logit) | |
| features.append(feats) | |
| return logits, features | |