# 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