# 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, } @torch.no_grad() 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, } @torch.no_grad() 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