Download Modules/rvq_codec.py from FashionFlora/SFlowTTS: direct link, hf CLI and curl.
- Browser
- Download file 28.2 kB
-
https://huggingface.co/FashionFlora/SFlowTTS/resolve/main/Modules/rvq_codec.py
- Command line
-
hf download hf://FashionFlora/SFlowTTS/Modules/rvq_codec.py
-
curl -L -o rvq_codec.py https://huggingface.co/FashionFlora/SFlowTTS/resolve/main/Modules/rvq_codec.py
28.2 kB
| # rvq_codec.py | |
| # ============================================================================== | |
| # Residual Vector Quantization (RVQ) Codec for TTS | |
| # | |
| # This module compresses pitch, energy, text embeddings, and style into | |
| # discrete tokens using RVQ-VAE technique. Similar to SoundStream/EnCodec | |
| # but adapted for TTS conditioning signals. | |
| # | |
| # Input: pitch[B,T], energy[B,T], text_emb[B,512,T], style[B,128] | |
| # Output: discrete tokens [B, num_quantizers, T_compressed] | |
| # ============================================================================== | |
| import math | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from torch.nn.utils import weight_norm, spectral_norm | |
| from einops import rearrange, repeat | |
| from typing import Tuple, Optional, List | |
| # ============================================================================== | |
| # Vector Quantizer with EMA updates | |
| # ============================================================================== | |
| class VectorQuantizerEMA(nn.Module): | |
| """ | |
| Improved VQ with Exponential Moving Average updates for codebook. | |
| Based on Neural Discrete Representation Learning (van den Oord et al.) | |
| """ | |
| def __init__( | |
| self, | |
| num_embeddings: int = 1024, | |
| embedding_dim: int = 256, | |
| commitment_cost: float = 0.25, | |
| decay: float = 0.99, | |
| epsilon: float = 1e-5, | |
| kmeans_init: bool = True, | |
| threshold_ema_dead_code: int = 2, | |
| ): | |
| super().__init__() | |
| self.num_embeddings = num_embeddings | |
| self.embedding_dim = embedding_dim | |
| self.commitment_cost = commitment_cost | |
| self.decay = decay | |
| self.epsilon = epsilon | |
| self.threshold_ema_dead_code = threshold_ema_dead_code | |
| self.kmeans_init = kmeans_init | |
| # Codebook | |
| embed = torch.randn(num_embeddings, embedding_dim) | |
| self.register_buffer("embed", embed) | |
| self.register_buffer("cluster_size", torch.zeros(num_embeddings)) | |
| self.register_buffer("embed_avg", embed.clone()) | |
| self.register_buffer("inited", torch.tensor([not kmeans_init])) | |
| def _init_embed(self, data: torch.Tensor): | |
| """Initialize codebook from first batch using k-means.""" | |
| if self.inited.item(): | |
| return | |
| # Flatten data for k-means | |
| flat = rearrange(data, "b d t -> (b t) d") | |
| if flat.shape[0] >= self.num_embeddings: | |
| # Random sample for init | |
| indices = torch.randperm(flat.shape[0])[:self.num_embeddings] | |
| embed = flat[indices] | |
| else: | |
| # Repeat if not enough samples | |
| repeats = (self.num_embeddings // flat.shape[0]) + 1 | |
| embed = flat.repeat(repeats, 1)[:self.num_embeddings] | |
| self.embed.data.copy_(embed) | |
| self.embed_avg.data.copy_(embed) | |
| self.cluster_size.data.fill_(1) | |
| self.inited.data.fill_(True) | |
| def _expire_codes(self, batch_samples: torch.Tensor): | |
| """Replace dead codes with random samples from batch.""" | |
| if self.threshold_ema_dead_code == 0: | |
| return | |
| dead_codes = self.cluster_size < self.threshold_ema_dead_code | |
| num_dead = dead_codes.sum().item() | |
| if num_dead == 0: | |
| return | |
| # Get random samples from batch | |
| flat = rearrange(batch_samples, "b d t -> (b t) d") | |
| indices = torch.randperm(flat.shape[0])[:num_dead] | |
| samples = flat[indices] | |
| # Replace dead codes | |
| self.embed.data[dead_codes] = samples | |
| self.embed_avg.data[dead_codes] = samples | |
| self.cluster_size.data[dead_codes] = 1 | |
| def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| """ | |
| Args: | |
| x: [B, D, T] input features | |
| Returns: | |
| quantized: [B, D, T] quantized features | |
| indices: [B, T] codebook indices | |
| loss: commitment + codebook loss | |
| """ | |
| B, D, T = x.shape | |
| # Initialize codebook on first forward | |
| if self.training: | |
| self._init_embed(x) | |
| # [B, D, T] -> [B, T, D] | |
| x_flat = rearrange(x, "b d t -> (b t) d") | |
| # Compute distances to codebook entries | |
| # ||x - e||^2 = ||x||^2 - 2*x*e + ||e||^2 | |
| distances = ( | |
| x_flat.pow(2).sum(dim=1, keepdim=True) | |
| - 2 * x_flat @ self.embed.t() | |
| + self.embed.pow(2).sum(dim=1, keepdim=True).t() | |
| ) | |
| # Get nearest codebook entries | |
| indices = distances.argmin(dim=1) # [(B*T)] | |
| # One-hot for EMA update | |
| encodings = F.one_hot(indices, self.num_embeddings).float() # [(B*T), K] | |
| # Quantize | |
| quantized = F.embedding(indices, self.embed) # [(B*T), D] | |
| # EMA codebook update | |
| if self.training: | |
| # Update cluster sizes | |
| self.cluster_size.data.mul_(self.decay).add_( | |
| encodings.sum(0), alpha=1 - self.decay | |
| ) | |
| # Update embedding averages | |
| embed_sum = encodings.t() @ x_flat # [K, D] | |
| self.embed_avg.data.mul_(self.decay).add_( | |
| embed_sum, alpha=1 - self.decay | |
| ) | |
| # Normalize | |
| n = self.cluster_size.sum() | |
| cluster_size = ( | |
| (self.cluster_size + self.epsilon) | |
| / (n + self.num_embeddings * self.epsilon) * n | |
| ) | |
| self.embed.data.copy_(self.embed_avg / cluster_size.unsqueeze(1)) | |
| # Expire dead codes | |
| self._expire_codes(x) | |
| # Commitment loss | |
| commitment_loss = F.mse_loss(quantized.detach(), x_flat) | |
| # Straight-through estimator | |
| quantized = x_flat + (quantized - x_flat).detach() | |
| # Reshape back | |
| quantized = rearrange(quantized, "(b t) d -> b d t", b=B, t=T) | |
| indices = rearrange(indices, "(b t) -> b t", b=B, t=T) | |
| loss = self.commitment_cost * commitment_loss | |
| return quantized, indices, loss | |
| def decode(self, indices: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Decode indices to embeddings. | |
| Args: | |
| indices: [B, T] or [B, T, num_quantizers] | |
| Returns: | |
| embeddings: [B, D, T] | |
| """ | |
| quantized = F.embedding(indices, self.embed) # [B, T, D] | |
| return rearrange(quantized, "b t d -> b d t") | |
| # ============================================================================== | |
| # Residual Vector Quantizer (RVQ) - Cascaded VQ layers | |
| # ============================================================================== | |
| class ResidualVectorQuantizer(nn.Module): | |
| """ | |
| Residual Vector Quantization with multiple codebooks. | |
| Each subsequent VQ quantizes the residual from previous. | |
| """ | |
| def __init__( | |
| self, | |
| num_quantizers: int = 8, | |
| num_embeddings: int = 1024, | |
| embedding_dim: int = 256, | |
| commitment_cost: float = 0.25, | |
| decay: float = 0.99, | |
| kmeans_init: bool = True, | |
| ): | |
| super().__init__() | |
| self.num_quantizers = num_quantizers | |
| self.embedding_dim = embedding_dim | |
| self.quantizers = nn.ModuleList([ | |
| VectorQuantizerEMA( | |
| num_embeddings=num_embeddings, | |
| embedding_dim=embedding_dim, | |
| commitment_cost=commitment_cost, | |
| decay=decay, | |
| kmeans_init=kmeans_init, | |
| ) | |
| for _ in range(num_quantizers) | |
| ]) | |
| def forward( | |
| self, | |
| x: torch.Tensor, | |
| n_quantizers: Optional[int] = None, | |
| ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| """ | |
| Args: | |
| x: [B, D, T] input features | |
| n_quantizers: number of quantizers to use (for training with dropout) | |
| Returns: | |
| quantized: [B, D, T] sum of all quantized residuals | |
| indices: [B, n_q, T] codebook indices for each quantizer | |
| loss: total commitment loss | |
| """ | |
| n_q = n_quantizers or self.num_quantizers | |
| residual = x | |
| quantized_out = torch.zeros_like(x) | |
| all_indices = [] | |
| total_loss = 0.0 | |
| for i in range(n_q): | |
| quantized, indices, loss = self.quantizers[i](residual) | |
| residual = residual - quantized.detach() # Detach to prevent gradient flow to previous | |
| quantized_out = quantized_out + quantized | |
| all_indices.append(indices) | |
| total_loss = total_loss + loss | |
| # Stack indices: [B, n_q, T] | |
| all_indices = torch.stack(all_indices, dim=1) | |
| return quantized_out, all_indices, total_loss / n_q | |
| def decode(self, indices: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Decode from indices. | |
| Args: | |
| indices: [B, n_q, T] indices for each quantizer | |
| Returns: | |
| quantized: [B, D, T] | |
| """ | |
| B, n_q, T = indices.shape | |
| quantized = torch.zeros(B, self.embedding_dim, T, device=indices.device) | |
| for i in range(n_q): | |
| quantized = quantized + self.quantizers[i].decode(indices[:, i]) | |
| return quantized | |
| # ============================================================================== | |
| # Encoder: Compresses inputs to latent space | |
| # ============================================================================== | |
| class ConvBlock(nn.Module): | |
| """Residual convolution block with snake activation.""" | |
| def __init__(self, dim: int, kernel_size: int = 7, dilation: int = 1): | |
| super().__init__() | |
| padding = (kernel_size - 1) * dilation // 2 | |
| self.conv = nn.Sequential( | |
| weight_norm(nn.Conv1d(dim, dim, kernel_size, dilation=dilation, padding=padding)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(dim, dim, 1)), | |
| ) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return x + self.conv(x) | |
| class EncoderBlock(nn.Module): | |
| """Downsampling encoder block.""" | |
| def __init__(self, dim_in: int, dim_out: int, stride: int = 2): | |
| super().__init__() | |
| self.residual = nn.Sequential( | |
| ConvBlock(dim_in, 7, dilation=1), | |
| ConvBlock(dim_in, 7, dilation=3), | |
| ConvBlock(dim_in, 7, dilation=9), | |
| ) | |
| self.downsample = weight_norm( | |
| nn.Conv1d(dim_in, dim_out, kernel_size=2*stride, stride=stride, padding=stride//2) | |
| ) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x = self.residual(x) | |
| x = self.downsample(x) | |
| return x | |
| class CodecEncoder(nn.Module): | |
| """ | |
| Encoder that fuses pitch, energy, text embeddings, and style | |
| into a compressed latent representation. | |
| Input: | |
| pitch: [B, T] | |
| energy: [B, T] | |
| text_emb: [B, 512, T] | |
| style: [B, 128] | |
| Output: | |
| latent: [B, latent_dim, T//compression_ratio] | |
| """ | |
| def __init__( | |
| self, | |
| text_dim: int = 512, | |
| style_dim: int = 128, | |
| latent_dim: int = 256, | |
| hidden_dim: int = 512, | |
| strides: List[int] = [2, 2, 2], # Total compression: 8x | |
| ): | |
| super().__init__() | |
| self.latent_dim = latent_dim | |
| self.compression_ratio = math.prod(strides) | |
| # Pitch/Energy projections (from scalar to hidden) | |
| self.pitch_proj = nn.Sequential( | |
| weight_norm(nn.Conv1d(1, 64, 7, padding=3)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(64, 128, 3, padding=1)), | |
| ) | |
| self.energy_proj = nn.Sequential( | |
| weight_norm(nn.Conv1d(1, 64, 7, padding=3)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(64, 128, 3, padding=1)), | |
| ) | |
| # Text embedding projection | |
| self.text_proj = weight_norm(nn.Conv1d(text_dim, hidden_dim - 256, 1)) | |
| # Style projection (broadcast over time) | |
| self.style_proj = nn.Linear(style_dim, hidden_dim) | |
| # Fusion and encoder | |
| self.pre_encoder = nn.Sequential( | |
| weight_norm(nn.Conv1d(hidden_dim * 2, hidden_dim, 7, padding=3)), | |
| nn.SiLU(), | |
| ) | |
| # Downsampling blocks | |
| self.encoder_blocks = nn.ModuleList() | |
| dim = hidden_dim | |
| for stride in strides: | |
| self.encoder_blocks.append( | |
| EncoderBlock(dim, min(dim * 2, 1024), stride=stride) | |
| ) | |
| dim = min(dim * 2, 1024) | |
| # Project to latent dim | |
| self.to_latent = nn.Sequential( | |
| ConvBlock(dim, 7), | |
| weight_norm(nn.Conv1d(dim, latent_dim, 1)), | |
| ) | |
| def forward( | |
| self, | |
| pitch: torch.Tensor, | |
| energy: torch.Tensor, | |
| text_emb: torch.Tensor, | |
| style: torch.Tensor, | |
| ) -> torch.Tensor: | |
| """ | |
| Args: | |
| pitch: [B, T] | |
| energy: [B, T] | |
| text_emb: [B, 512, T] | |
| style: [B, 128] | |
| Returns: | |
| latent: [B, latent_dim, T//compression_ratio] | |
| """ | |
| B, T = pitch.shape | |
| # Process pitch and energy | |
| pitch_feat = self.pitch_proj(pitch.unsqueeze(1)) # [B, 128, T] | |
| energy_feat = self.energy_proj(energy.unsqueeze(1)) # [B, 128, T] | |
| # Process text | |
| text_feat = self.text_proj(text_emb) # [B, hidden-256, T] | |
| # Concatenate pitch, energy, text | |
| cond_feat = torch.cat([pitch_feat, energy_feat, text_feat], dim=1) # [B, hidden, T] | |
| # Broadcast style over time and add | |
| style_feat = self.style_proj(style) # [B, hidden] | |
| style_feat = style_feat.unsqueeze(-1).expand(-1, -1, T) # [B, hidden, T] | |
| # Fuse all features | |
| x = torch.cat([cond_feat, style_feat], dim=1) # [B, hidden*2, T] | |
| x = self.pre_encoder(x) | |
| # Encode | |
| for block in self.encoder_blocks: | |
| x = block(x) | |
| # Project to latent | |
| latent = self.to_latent(x) | |
| return latent | |
| # ============================================================================== | |
| # Decoder: Reconstructs from quantized latents | |
| # ============================================================================== | |
| class DecoderBlock(nn.Module): | |
| """Upsampling decoder block with style conditioning.""" | |
| def __init__(self, dim_in: int, dim_out: int, style_dim: int = 128, stride: int = 2): | |
| super().__init__() | |
| self.upsample = weight_norm( | |
| nn.ConvTranspose1d(dim_in, dim_out, kernel_size=2*stride, stride=stride, padding=stride//2) | |
| ) | |
| self.residual = nn.Sequential( | |
| ConvBlock(dim_out, 7, dilation=1), | |
| ConvBlock(dim_out, 7, dilation=3), | |
| ConvBlock(dim_out, 7, dilation=9), | |
| ) | |
| # Style conditioning via FiLM | |
| self.style_proj = nn.Linear(style_dim, dim_out * 2) | |
| def forward(self, x: torch.Tensor, style: torch.Tensor) -> torch.Tensor: | |
| x = self.upsample(x) | |
| # FiLM conditioning | |
| style_params = self.style_proj(style) # [B, dim*2] | |
| gamma, beta = style_params.chunk(2, dim=-1) | |
| gamma = gamma.unsqueeze(-1) # [B, dim, 1] | |
| beta = beta.unsqueeze(-1) | |
| x = x * (1 + gamma) + beta | |
| x = self.residual(x) | |
| return x | |
| class CodecDecoder(nn.Module): | |
| """ | |
| Decoder that reconstructs conditioning signals from quantized latents. | |
| Input: | |
| latent: [B, latent_dim, T_compressed] | |
| style: [B, 128] | |
| Output: | |
| pitch: [B, T] | |
| energy: [B, T] | |
| text_emb: [B, 512, T] | |
| """ | |
| def __init__( | |
| self, | |
| text_dim: int = 512, | |
| style_dim: int = 128, | |
| latent_dim: int = 256, | |
| hidden_dim: int = 512, | |
| strides: List[int] = [2, 2, 2], | |
| ): | |
| super().__init__() | |
| self.compression_ratio = math.prod(strides) | |
| # From latent to decoder | |
| self.from_latent = nn.Sequential( | |
| weight_norm(nn.Conv1d(latent_dim, 1024, 1)), | |
| nn.SiLU(), | |
| ) | |
| # Upsampling blocks | |
| self.decoder_blocks = nn.ModuleList() | |
| dims = [1024] | |
| dim = 1024 | |
| for stride in strides: | |
| dim_out = max(dim // 2, hidden_dim) | |
| self.decoder_blocks.append( | |
| DecoderBlock(dim, dim_out, style_dim=style_dim, stride=stride) | |
| ) | |
| dim = dim_out | |
| dims.append(dim_out) | |
| # Output projections | |
| self.pitch_head = nn.Sequential( | |
| weight_norm(nn.Conv1d(hidden_dim, 128, 3, padding=1)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(128, 1, 3, padding=1)), | |
| ) | |
| self.energy_head = nn.Sequential( | |
| weight_norm(nn.Conv1d(hidden_dim, 128, 3, padding=1)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(128, 1, 3, padding=1)), | |
| ) | |
| self.text_head = nn.Sequential( | |
| weight_norm(nn.Conv1d(hidden_dim, hidden_dim, 3, padding=1)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(hidden_dim, text_dim, 1)), | |
| ) | |
| def forward( | |
| self, | |
| latent: torch.Tensor, | |
| style: torch.Tensor, | |
| target_len: Optional[int] = None, | |
| ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| """ | |
| Args: | |
| latent: [B, latent_dim, T_compressed] | |
| style: [B, 128] | |
| target_len: target output length (optional) | |
| Returns: | |
| pitch: [B, T] | |
| energy: [B, T] | |
| text_emb: [B, 512, T] | |
| """ | |
| x = self.from_latent(latent) | |
| for block in self.decoder_blocks: | |
| x = block(x, style) | |
| # Adjust length if needed | |
| if target_len is not None and x.shape[-1] != target_len: | |
| x = F.interpolate(x, size=target_len, mode="linear", align_corners=False) | |
| # Generate outputs | |
| pitch = self.pitch_head(x).squeeze(1) # [B, T] | |
| energy = self.energy_head(x).squeeze(1) # [B, T] | |
| text_emb = self.text_head(x) # [B, 512, T] | |
| return pitch, energy, text_emb | |
| # ============================================================================== | |
| # Full Codec Model | |
| # ============================================================================== | |
| class TTSCodec(nn.Module): | |
| """ | |
| Full TTS Codec: Encoder -> RVQ -> Decoder | |
| Compresses pitch, energy, text embeddings, and style into discrete tokens. | |
| """ | |
| def __init__( | |
| self, | |
| # Dimensions | |
| text_dim: int = 512, | |
| style_dim: int = 128, | |
| latent_dim: int = 256, | |
| hidden_dim: int = 512, | |
| # Compression | |
| strides: List[int] = [2, 2, 2], | |
| # RVQ | |
| num_quantizers: int = 8, | |
| codebook_size: int = 1024, | |
| commitment_cost: float = 0.25, | |
| ): | |
| super().__init__() | |
| self.text_dim = text_dim | |
| self.style_dim = style_dim | |
| self.latent_dim = latent_dim | |
| self.compression_ratio = math.prod(strides) | |
| self.num_quantizers = num_quantizers | |
| # Encoder | |
| self.encoder = CodecEncoder( | |
| text_dim=text_dim, | |
| style_dim=style_dim, | |
| latent_dim=latent_dim, | |
| hidden_dim=hidden_dim, | |
| strides=strides, | |
| ) | |
| # RVQ | |
| self.quantizer = ResidualVectorQuantizer( | |
| num_quantizers=num_quantizers, | |
| num_embeddings=codebook_size, | |
| embedding_dim=latent_dim, | |
| commitment_cost=commitment_cost, | |
| ) | |
| # Decoder | |
| self.decoder = CodecDecoder( | |
| text_dim=text_dim, | |
| style_dim=style_dim, | |
| latent_dim=latent_dim, | |
| hidden_dim=hidden_dim, | |
| strides=strides, | |
| ) | |
| def encode( | |
| self, | |
| pitch: torch.Tensor, | |
| energy: torch.Tensor, | |
| text_emb: torch.Tensor, | |
| style: torch.Tensor, | |
| ) -> torch.Tensor: | |
| """Encode inputs to continuous latent.""" | |
| return self.encoder(pitch, energy, text_emb, style) | |
| def quantize( | |
| self, | |
| latent: torch.Tensor, | |
| n_quantizers: Optional[int] = None, | |
| ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| """Quantize latent to discrete tokens.""" | |
| return self.quantizer(latent, n_quantizers) | |
| def decode( | |
| self, | |
| latent: torch.Tensor, | |
| style: torch.Tensor, | |
| target_len: Optional[int] = None, | |
| ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| """Decode from continuous latent.""" | |
| return self.decoder(latent, style, target_len) | |
| def decode_from_tokens( | |
| self, | |
| tokens: torch.Tensor, | |
| style: torch.Tensor, | |
| target_len: Optional[int] = None, | |
| ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| """ | |
| Decode directly from discrete tokens. | |
| Args: | |
| tokens: [B, num_quantizers, T_compressed] | |
| style: [B, 128] | |
| """ | |
| latent = self.quantizer.decode(tokens) | |
| return self.decode(latent, style, target_len) | |
| def forward( | |
| self, | |
| pitch: torch.Tensor, | |
| energy: torch.Tensor, | |
| text_emb: torch.Tensor, | |
| style: torch.Tensor, | |
| n_quantizers: Optional[int] = None, | |
| ) -> dict: | |
| """ | |
| Full forward pass: encode -> quantize -> decode | |
| Returns dict with: | |
| - pitch_rec: reconstructed pitch | |
| - energy_rec: reconstructed energy | |
| - text_emb_rec: reconstructed text embeddings | |
| - tokens: discrete tokens [B, n_q, T_compressed] | |
| - quantized: quantized latent | |
| - commitment_loss: RVQ commitment loss | |
| """ | |
| T = pitch.shape[-1] | |
| # Encode | |
| latent = self.encode(pitch, energy, text_emb, style) | |
| # Quantize | |
| quantized, tokens, commitment_loss = self.quantize(latent, n_quantizers) | |
| # Decode | |
| pitch_rec, energy_rec, text_emb_rec = self.decode(quantized, style, target_len=T) | |
| return { | |
| "pitch_rec": pitch_rec, | |
| "energy_rec": energy_rec, | |
| "text_emb_rec": text_emb_rec, | |
| "tokens": tokens, | |
| "latent": latent, | |
| "quantized": quantized, | |
| "commitment_loss": commitment_loss, | |
| } | |
| def tokenize( | |
| self, | |
| pitch: torch.Tensor, | |
| energy: torch.Tensor, | |
| text_emb: torch.Tensor, | |
| style: torch.Tensor, | |
| ) -> torch.Tensor: | |
| """ | |
| Get discrete tokens for inputs (inference mode). | |
| Returns: tokens [B, num_quantizers, T_compressed] | |
| """ | |
| latent = self.encode(pitch, energy, text_emb, style) | |
| _, tokens, _ = self.quantize(latent) | |
| return tokens | |
| # ============================================================================== | |
| # Combined Codec + Vocoder for end-to-end training | |
| # ============================================================================== | |
| class CodecVocoder(nn.Module): | |
| """ | |
| Combined Codec and Vocoder that: | |
| 1. Encodes pitch/energy/text/style into discrete tokens | |
| 2. Decodes tokens back to conditioning signals | |
| 3. Generates waveform from reconstructed conditioning + style | |
| This enables end-to-end training with waveform reconstruction loss. | |
| """ | |
| def __init__( | |
| self, | |
| codec: TTSCodec, | |
| vocoder: nn.Module, # Your ringformer decoder | |
| ): | |
| super().__init__() | |
| self.codec = codec | |
| self.vocoder = vocoder | |
| def forward( | |
| self, | |
| pitch: torch.Tensor, | |
| energy: torch.Tensor, | |
| text_emb: torch.Tensor, | |
| style: torch.Tensor, | |
| n_quantizers: Optional[int] = None, | |
| ) -> dict: | |
| """ | |
| Full pipeline: inputs -> tokens -> reconstruction -> waveform | |
| """ | |
| # Codec forward | |
| codec_out = self.codec(pitch, energy, text_emb, style, n_quantizers) | |
| # Generate waveform from reconstructed conditions | |
| wav_rec, mag, phase = self.vocoder( | |
| codec_out["text_emb_rec"], | |
| codec_out["pitch_rec"], | |
| codec_out["energy_rec"], | |
| style, | |
| ) | |
| # Also generate from original for comparison | |
| wav_orig, _, _ = self.vocoder(text_emb, pitch, energy, style) | |
| return { | |
| **codec_out, | |
| "wav_rec": wav_rec, | |
| "wav_orig": wav_orig, | |
| "mag": mag, | |
| "phase": phase, | |
| } | |
| def generate_from_tokens( | |
| self, | |
| tokens: torch.Tensor, | |
| style: torch.Tensor, | |
| target_len: int, | |
| ) -> torch.Tensor: | |
| """ | |
| Generate waveform directly from discrete tokens. | |
| Args: | |
| tokens: [B, num_quantizers, T_compressed] | |
| style: [B, 128] | |
| target_len: target sequence length | |
| """ | |
| pitch, energy, text_emb = self.codec.decode_from_tokens(tokens, style, target_len) | |
| wav, _, _ = self.vocoder(text_emb, pitch, energy, style) | |
| return wav | |
| # ============================================================================== | |
| # Finite Scalar Quantization alternative (FSQ) for reference | |
| # ============================================================================== | |
| class FiniteScalarQuantizer(nn.Module): | |
| """ | |
| Finite Scalar Quantization (FSQ) - simpler alternative to VQ. | |
| Maps continuous values to a fixed number of levels per dimension. | |
| """ | |
| def __init__(self, levels: List[int], dim: int = 256): | |
| super().__init__() | |
| self.levels = levels | |
| self.num_levels = levels | |
| self.dim = dim | |
| # Number of dimensions that will be quantized | |
| self.n_codes = len(levels) | |
| # Projection to quantized dimensions | |
| if dim != self.n_codes: | |
| self.proj_in = nn.Linear(dim, self.n_codes) | |
| self.proj_out = nn.Linear(self.n_codes, dim) | |
| else: | |
| self.proj_in = nn.Identity() | |
| self.proj_out = nn.Identity() | |
| # Register levels as buffer | |
| self.register_buffer("_levels", torch.tensor(levels)) | |
| def _round_ste(self, x: torch.Tensor) -> torch.Tensor: | |
| """Round with straight-through estimator.""" | |
| return x + (x.round() - x).detach() | |
| def _quantize(self, x: torch.Tensor) -> torch.Tensor: | |
| """Quantize to discrete levels.""" | |
| # Scale to [-1, 1] then to [0, L-1] | |
| x = torch.tanh(x) # [-1, 1] | |
| # Per-dimension levels | |
| half_levels = (self._levels - 1) / 2 | |
| x = x * half_levels | |
| x = self._round_ste(x) | |
| x = x / half_levels # Back to [-1, 1] | |
| return x | |
| def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: | |
| """ | |
| Args: | |
| x: [B, D, T] input features | |
| Returns: | |
| quantized: [B, D, T] | |
| indices: [B, T] flattened indices (for logging) | |
| """ | |
| # [B, D, T] -> [B, T, D] | |
| x = rearrange(x, "b d t -> b t d") | |
| x = self.proj_in(x) | |
| x_q = self._quantize(x) | |
| x_out = self.proj_out(x_q) | |
| # Compute indices (for analysis) | |
| half_levels = (self._levels - 1) / 2 | |
| indices = ((x_q * half_levels) + half_levels).long() | |
| # Flatten multi-dim indices to single index | |
| multipliers = torch.cumprod( | |
| torch.cat([torch.ones(1, device=x.device), self._levels[:-1].float()]), | |
| dim=0 | |
| ).long() | |
| flat_indices = (indices * multipliers).sum(dim=-1) # [B, T] | |
| x_out = rearrange(x_out, "b t d -> b d t") | |
| return x_out, flat_indices | |