|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 *
|
|
|
|
|
|
|
|
|
| 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:
|
|
|
| T = x.size(1)
|
| return x + self.pe[:, :T, :]
|
|
|
|
|
|
|
|
|
|
|
| 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,
|
|
|
| num_languages: int = 0,
|
| max_speakers_per_language: int = 0,
|
| 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
|
|
|
|
|
| self.num_languages = num_languages
|
| self.max_speakers_per_language = max_speakers_per_language
|
|
|
|
|
| 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
|
|
|
|
|
| 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
|
|
|
|
|
| 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
|
|
|
|
|
| self.cls = nn.Parameter(torch.randn(1, 1, self.total_style_dim) * 0.02)
|
|
|
|
|
| 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,
|
| embedding_max_length=max_len,
|
| dropout=dropout,
|
| attn_dropout=attn_dropout,
|
| ff_dropout=ff_dropout,
|
| norm_type=norm_type,
|
| )
|
|
|
|
|
| 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
|
|
|
|
|
| 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 = []
|
|
|
|
|
| if self.lang_embed is not None:
|
| assert language_ids is not None
|
| lang_vec = self.lang_embed(language_ids)
|
| lang = lang_vec.unsqueeze(1).expand(-1, T, -1)
|
| extras.append(lang)
|
|
|
|
|
| if self.spk_embed is not None:
|
| assert speaker_ids is not None and language_ids is not None
|
|
|
|
|
| spk_global = self._compute_global_speaker_ids(speaker_ids, language_ids)
|
|
|
| spk_vec = self.spk_embed(spk_global)
|
| spk = spk_vec.unsqueeze(1).expand(-1, T, -1)
|
| extras.append(spk)
|
|
|
| if extras:
|
|
|
| tokens = torch.cat([tokens] + extras, dim=-1)
|
| tokens = self.cond_proj(tokens)
|
|
|
|
|
| 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)
|
|
|
|
|
| tokens, mask = self._fuse_condition(
|
| cond_tokens, cond_mask, speaker_ids, language_ids
|
| )
|
|
|
|
|
| cls = self.cls.expand(B, 1, -1)
|
|
|
|
|
|
|
|
|
| out = self.transformer.forward(
|
| cls,
|
| None,
|
| embedding_mask_proba=self.embedding_mask_proba,
|
| embedding=tokens,
|
| embedding_scale=1.0,
|
| )
|
|
|
| z_hat = out.squeeze(1)
|
| 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
|
|
|
|
|
| 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
|
|
|
|
|
| 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)
|
| 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)
|
| 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
|
| )
|
| spk_vec = self.spk_embed(spk_global)
|
| 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
|
| )
|
| 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
|
|
|
|
|