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