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