SFlowTTS / Modules /durations_rl.py
FashionFlora's picture
Upload full repo excluding dump_40, dump_100, precomputed_tokens, precomputed_data
fb0011a verified
Raw History Blame Contribute Delete
37.8 kB
# models/reinforce_aligner.py
# Core encoder + phoneme-wise (α=2) Reinforce-Aligner with Gaussian upsampling
# Modified to compute rewards from GT durations instead of mel spectrogram loss.
from typing import Callable, Dict, Optional, Tuple
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
def rms_norm_fn(x: torch.Tensor, eps: float = 1e-5) -> torch.Tensor:
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + eps)
class AdaRMSNorm(nn.Module):
"""
AdaRMSNorm: (1 + gamma) * rms_norm(x) + beta.
Gamma i Beta generowane dynamicznie z wektora stylu.
"""
def __init__(self, style_dim: int, channels: int, eps: float = 1e-5) -> None:
super().__init__()
self.fc = nn.Linear(style_dim, channels * 2)
self.eps = eps
# Inicjalizacja "identity" - gamma i beta bliskie 0 na start
with torch.no_grad():
self.fc.weight.zero_()
self.fc.bias.zero_()
def forward(self, x: torch.Tensor, style: torch.Tensor) -> torch.Tensor:
# x: [B, T, C], style: [B, S]
gb = self.fc(style) # [B, 2C]
gamma, beta = gb.chunk(2, dim=-1)
# Rozszerzenie wymiarów do [B, 1, C] dla broadcastingu po czasie
gamma = gamma.unsqueeze(1).to(dtype=x.dtype)
beta = beta.unsqueeze(1).to(dtype=x.dtype)
x_norm = x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
return (1.0 + gamma) * x_norm + beta
class SwiGLUFFN(nn.Module):
def __init__(self, dim: int, expansion: float = 2.0, dropout: float = 0.0) -> None:
super().__init__()
hidden = int(dim * expansion)
self.fc = nn.Linear(dim, hidden * 2)
self.proj = nn.Linear(hidden, dim)
self.dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
a, b = self.fc(x).chunk(2, dim=-1)
x = F.silu(b) * a
x = self.proj(self.dropout(x))
return x
# -------------------- RoPE & Attention --------------------
class RotaryEmbedding(nn.Module):
def __init__(self, dim: int, base: float = 10000.0) -> None:
super().__init__()
if dim % 2 != 0:
raise ValueError("RoPE head_dim must be even.")
inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))
self.register_buffer("inv_freq", inv_freq, persistent=False)
def get_cos_sin(
self, seq_len: int, device: torch.device, dtype: torch.dtype
) -> Tuple[torch.Tensor, torch.Tensor]:
t = torch.arange(seq_len, device=device, dtype=dtype)
freqs = torch.outer(t, self.inv_freq.to(dtype))
emb = torch.cat([freqs, freqs], dim=-1)
cos = emb.cos().unsqueeze(0).unsqueeze(0)
sin = emb.sin().unsqueeze(0).unsqueeze(0)
return cos, sin
def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
x_even = x[..., ::2]
x_odd = x[..., 1::2]
x_rot = torch.stack([-x_odd, x_even], dim=-1).flatten(-2)
return (x * cos) + (x_rot * sin)
class MultiheadSelfAttentionRoPE(nn.Module):
def __init__(
self,
embed_dim: int,
num_heads: int,
dropout: float = 0.0,
use_rope: bool = True,
rope_base: float = 10000.0,
) -> None:
super().__init__()
self.embed_dim = embed_dim
self.num_heads = num_heads
self.head_dim = embed_dim // num_heads
self.scale = 1.0 / math.sqrt(self.head_dim)
self.qkv = nn.Linear(embed_dim, embed_dim * 3)
self.out_proj = nn.Linear(embed_dim, embed_dim)
self.attn_dropout = nn.Dropout(dropout)
self.use_rope = use_rope
self.rope = RotaryEmbedding(self.head_dim, base=rope_base) if use_rope else None
def forward(
self, x: torch.Tensor, key_padding_mask: Optional[torch.Tensor] = None
) -> torch.Tensor:
B, T, C = x.shape
qkv = self.qkv(x)
q, k, v = qkv.chunk(3, dim=-1)
# [B, T, 3C] -> [B, H, T, D]
q = q.view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
k = k.view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
v = v.view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
if self.use_rope:
cos, sin = self.rope.get_cos_sin(T, x.device, x.dtype)
q = apply_rope(q, cos, sin)
k = apply_rope(k, cos, sin)
attn_scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale
if key_padding_mask is not None:
# Mask format: [B, T] -> [B, 1, 1, T] for broadcasting
mask = key_padding_mask.unsqueeze(1).unsqueeze(2)
attn_scores = attn_scores.masked_fill(mask, torch.finfo(attn_scores.dtype).min)
attn_weights = F.softmax(attn_scores, dim=-1)
attn_weights = self.attn_dropout(attn_weights)
out = torch.matmul(attn_weights, v)
out = out.transpose(1, 2).contiguous().view(B, T, C)
return self.out_proj(out)
# -------------------- Reasoning Components --------------------
class ReasoningBlock(nn.Module):
def __init__(
self,
dim: int,
num_heads: int,
expansion: float,
style_dim: int,
attn_dropout: float = 0.0,
resid_dropout: float = 0.0,
rms_eps: float = 1e-5,
use_rope: bool = True,
rope_base: float = 10000.0,
) -> None:
super().__init__()
self.attn = MultiheadSelfAttentionRoPE(
embed_dim=dim,
num_heads=num_heads,
dropout=attn_dropout,
use_rope=use_rope,
rope_base=rope_base,
)
self.ff = SwiGLUFFN(dim=dim, expansion=expansion, dropout=resid_dropout)
self.resid_dropout = nn.Dropout(resid_dropout)
# Stylizacja poprzez AdaRMSNorm
self.ada1 = AdaRMSNorm(style_dim=style_dim, channels=dim, eps=rms_eps)
self.ada2 = AdaRMSNorm(style_dim=style_dim, channels=dim, eps=rms_eps)
def forward(
self,
hidden_states: torch.Tensor,
style: torch.Tensor,
key_padding_mask: Optional[torch.Tensor],
) -> torch.Tensor:
# 1. Attention Block z AdaNorm
attn_out = self.attn(hidden_states, key_padding_mask=key_padding_mask)
hidden_states = self.ada1(
hidden_states + self.resid_dropout(attn_out), style
)
# 2. FFN Block z AdaNorm
mlp_out = self.ff(hidden_states)
hidden_states = self.ada2(
hidden_states + self.resid_dropout(mlp_out), style
)
return hidden_states
class ReasoningModule(nn.Module):
def __init__(self, layers: nn.ModuleList) -> None:
super().__init__()
self.layers = layers
def forward(
self,
hidden_states: torch.Tensor,
input_injection: torch.Tensor,
style: torch.Tensor,
key_padding_mask: Optional[torch.Tensor],
) -> torch.Tensor:
# Standardowa architektura rekurencyjna często dodaje wejście w każdym kroku (Input Injection)
# zapobiega to "zapominaniu" promptu/kontekstu.
hidden_states = hidden_states + input_injection
for layer in self.layers:
hidden_states = layer(
hidden_states=hidden_states,
style=style,
key_padding_mask=key_padding_mask,
)
return hidden_states
def sinusoidal_positional_encoding(
seq_len: int, dim: int, device: torch.device, dtype: torch.dtype
) -> torch.Tensor:
pos = torch.arange(seq_len, device=device, dtype=dtype).unsqueeze(1)
i = torch.arange(dim, device=device, dtype=dtype).unsqueeze(0)
rates = torch.pow(10000.0, -(2 * torch.div(i, 2, rounding_mode="floor")) / dim)
angles = pos * rates
pe = torch.zeros((seq_len, dim), device=device, dtype=dtype)
pe[:, 0::2] = torch.sin(angles[:, 0::2])
pe[:, 1::2] = torch.cos(angles[:, 1::2])
return pe.unsqueeze(0)
# -------------------- FULL HRM ENCODER --------------------
class HRMDurationEncoder(nn.Module):
"""
Pełna implementacja HRM (Hierarchical Reasoning Model).
Zmiany względem wersji uproszczonej:
1. Usunięto 'torch.no_grad()' w pętli -> Pełne Backpropagation Through Time (BPTT).
2. Dodano Step Embeddings dla poziomów H i L -> Model wie, w której iteracji myślenia jest.
3. Zachowano architekturę System 1 / System 2 (Hierarchical Recurrence).
"""
def __init__(
self,
sty_dim: int,
d_model: int,
nlayers: int,
max_dur: int,
dropout: float = 0.1,
num_heads: int = 8,
expansion: float = 2.0,
H_cycles: int = 1, # Liczba cykli "wolnych" (High-level)
L_cycles: int = 2, # Liczba cykli "szybkich" (Low-level) wewnątrz jednego H
rms_norm_eps: float = 1e-5,
use_posenc: bool = False,
use_rope: bool = True,
rope_base: float = 10000.0,
vocab_size: Optional[int] = None,
pad_token_id: Optional[int] = None,
token_emb_dim: Optional[int] = None,
token_emb_dropout: float = 0.0,
token_embedding: Optional[nn.Embedding] = None,
token_proj: Optional[nn.Module] = None,
) -> None:
super().__init__()
self.d_model = d_model
self.sty_dim = sty_dim
self.H_cycles = H_cycles
self.L_cycles = L_cycles
self.use_posenc = use_posenc
# --- Embeddingi ---
self.pad_token_id = pad_token_id
if token_embedding is not None:
self.token_embedding = token_embedding
emb_dim = token_embedding.embedding_dim
elif vocab_size is not None:
emb_dim = token_emb_dim or d_model
self.token_embedding = nn.Embedding(vocab_size, emb_dim, padding_idx=pad_token_id)
else:
self.token_embedding = None
emb_dim = d_model # Zakładamy wejście continuous
self.token_emb_dropout = nn.Dropout(token_emb_dropout)
# Projekcja wejścia (jeśli dim się nie zgadza)
if emb_dim != d_model:
self.token_proj = nn.Linear(emb_dim, d_model)
else:
self.token_proj = nn.Identity()
# Input Projection: łączy embedding treści z embeddingiem stylu
self.in_proj = nn.Linear(d_model + sty_dim, d_model)
self.dropout = nn.Dropout(dropout)
# --- HRM Core ---
# Inicjalizacja stanów ukrytych (parametry uczące się)
self.H_init = nn.Parameter(torch.zeros(1, 1, d_model))
self.L_init = nn.Parameter(torch.zeros(1, 1, d_model))
nn.init.trunc_normal_(self.H_init, std=0.02)
nn.init.trunc_normal_(self.L_init, std=0.02)
# Step Embeddings - KLUCZOWE dla pełnego HRM
# Pozwalają modelowi odróżnić pierwszy krok rozumowania od ostatniego
self.h_step_emb = nn.Embedding(H_cycles, d_model)
self.l_step_emb = nn.Embedding(L_cycles, d_model)
# Moduły rezonowania (Recurrent Bodies)
# Zauważ: ReasoningModule zawiera listę ReasoningBlocks (Stack Transformerów)
self.H_level = ReasoningModule(nn.ModuleList([
ReasoningBlock(
dim=d_model, num_heads=num_heads, expansion=expansion,
style_dim=sty_dim, attn_dropout=dropout, resid_dropout=dropout,
rms_eps=rms_norm_eps, use_rope=use_rope, rope_base=rope_base
) for _ in range(nlayers)
]))
self.L_level = ReasoningModule(nn.ModuleList([
ReasoningBlock(
dim=d_model, num_heads=num_heads, expansion=expansion,
style_dim=sty_dim, attn_dropout=dropout, resid_dropout=dropout,
rms_eps=rms_norm_eps, use_rope=use_rope, rope_base=rope_base
) for _ in range(nlayers)
]))
# --- Output ---
self.dur_head = nn.Linear(d_model, max_dur)
def _prepare_input(self, x, m, device):
"""Obsługa konwersji ID -> Embedding i maskowania."""
# Sprawdzenie czy x to [B, T] (tokeny) czy [B, T, C] (embeddingi)
if x.dtype in [torch.long, torch.int] and x.dim() == 2:
if self.token_embedding is None:
raise ValueError("Otrzymano token IDs, ale brak token_embedding.")
# Automatyczna detekcja maski jeśli nie podana
if m is None and self.pad_token_id is not None:
m = (x == self.pad_token_id)
x = self.token_embedding(x)
x = self.token_emb_dropout(x)
x = self.token_proj(x)
else:
# Zakładamy continuous [B, T, C]
if x.dim() == 3 and x.size(1) == self.d_model: # [B, C, T] -> [B, T, C]
x = x.transpose(1, 2)
return x, m
def forward(
self,
x: torch.Tensor,
style: torch.Tensor,
text_lengths: Optional[torch.Tensor] = None, # Opcjonalne dla kompatybilności
m: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""
x: Token IDs [B, T] lub Embeddings [B, T, C]
style: Style vector [B, S]
m: Padding mask [B, T] (True = padding)
"""
device = x.device
B = x.shape[0]
# 1. Przygotowanie wejścia
x, m = self._prepare_input(x, m, device)
T = x.shape[1]
# 2. Kondycjonowanie wejścia stylem (Time-Distributed)
s_expanded = style.unsqueeze(1).expand(B, T, -1)
x_in = torch.cat([x, s_expanded], dim=-1)
x_in = self.in_proj(x_in)
x_in = self.dropout(x_in)
# Opcjonalne klasyczne kodowanie pozycyjne (zwykle zbędne przy RoPE)
if self.use_posenc:
pe = sinusoidal_positional_encoding(T, self.d_model, device, x.dtype)
x_in = x_in + pe
# 3. Inicjalizacja stanów HRM
# [1, 1, C] -> [B, T, C]
z_H = self.H_init.to(dtype=x.dtype).expand(B, T, -1)
z_L = self.L_init.to(dtype=x.dtype).expand(B, T, -1)
# Maska dla attention
key_padding_mask = m.to(device) if m is not None else None
# 4. Główna pętla Hierarchical Reasoning (Full BPTT)
# H (Slow) - zewnętrzna pętla
for h in range(self.H_cycles):
# Dodanie embeddingu kroku dla H
h_step = torch.tensor(h, device=device)
h_emb = self.h_step_emb(h_step).view(1, 1, -1)
# Wstrzyknięcie informacji o kroku do stanu H
z_H_curr = z_H + h_emb
# L (Fast) - wewnętrzna pętla
for l in range(self.L_cycles):
l_step = torch.tensor(l, device=device)
l_emb = self.l_step_emb(l_step).view(1, 1, -1)
# Wejście do L to: aktualny stan L + kontekst z H + oryginalny input
# To pozwala L "obrabiać" dane w kontekście wyższego poziomu
current_input_L = z_H_curr + x_in + l_emb
z_L = self.L_level(
hidden_states=z_L,
input_injection=current_input_L,
style=style,
key_padding_mask=key_padding_mask
)
# Opcjonalne czyszczenie paddingu po każdym kroku (dla stabilności)
if key_padding_mask is not None:
z_L = z_L.masked_fill(key_padding_mask.unsqueeze(-1), 0.0)
# Po zakończeniu cykli L, aktualizujemy poziom H
# Wejście do H to: aktualny stan H + wynik przemyśleń L
current_input_H = z_L # + x_in (opcjonalnie, ale L już to widziało)
# Zauważ: H_level w pętli zewnętrznej widzi skumulowany efekt L
if h < self.H_cycles: # Zawsze wykonujemy, warunek logiczny dla jasności
z_H = self.H_level(
hidden_states=z_H,
input_injection=current_input_H,
style=style,
key_padding_mask=key_padding_mask
)
if key_padding_mask is not None:
z_H = z_H.masked_fill(key_padding_mask.unsqueeze(-1), 0.0)
# 5. Zwracamy stan H (high-level representation) lub predykcję trwania
# Oryginalny kod zwracał z_H, więc zachowujemy to (głowicę można użyć na zewnątrz)
return z_H
def length_to_mask_second(lengths: torch.Tensor, max_len: Optional[int] = None) -> torch.Tensor:
# lengths: [B] (int)
max_len = int(max_len or lengths.max().item())
rng = torch.arange(max_len, device=lengths.device).unsqueeze(0) # [1, T]
mask = rng < lengths.unsqueeze(1) # [B, T]
return mask
class ResBlock1D(nn.Module):
def __init__(self, channels: int, kernel_size: int, dilations=(1, 3, 5), dropout=0.0):
super().__init__()
layers = []
for d in dilations:
pad = (kernel_size - 1) // 2 * d
layers += [
nn.LeakyReLU(0.2),
nn.Conv1d(channels, channels, kernel_size, padding=pad, dilation=d),
nn.Dropout(dropout),
]
self.body = nn.Sequential(*layers)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return x + self.body(x)
class MRFEncoder1D(nn.Module):
"""
Multi-Receptive Field stack (inspired by HiFi-GAN MRF).
Input: [B, D_in, N_tokens]
Output: [B, D, N_tokens] with D=out_channels
"""
def __init__(
self,
in_channels: int,
out_channels: int = 512,
num_blocks: int = 4,
kernels=(3, 5, 7),
dilations=(1, 3, 5),
dropout: float = 0.0,
):
super().__init__()
self.in_proj = nn.Conv1d(in_channels, out_channels, kernel_size=1)
blocks = []
for _ in range(num_blocks):
# parallel branches then sum
branches = nn.ModuleList([ResBlock1D(out_channels, k, dilations, dropout) for k in kernels])
blocks.append(branches)
self.blocks = nn.ModuleList(blocks)
def forward(self, x: torch.Tensor, token_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
# x: [B, Din, N]
x = self.in_proj(x) # [B, D, N]
for branches in self.blocks:
y = 0
for b in branches:
y = y + b(x)
x = y / len(branches)
if token_mask is not None:
x = x * token_mask.unsqueeze(1).float()
return x
class DurationPredictor(nn.Module):
"""
1D conv duration predictor over token axis.
Input: [B, D, N]
Output: durations_int [B, N] (positive ints), durations_raw [B, N] (float pre-round)
"""
def __init__(self, channels: int = 512, dropout: float = 0.1):
super().__init__()
self.conv1 = nn.Conv1d(channels, channels, kernel_size=3, padding=1)
self.ln1 = nn.LayerNorm(channels)
self.conv2 = nn.Conv1d(channels, channels, kernel_size=3, padding=1)
self.ln2 = nn.LayerNorm(channels)
self.proj = nn.Conv1d(channels, 1, kernel_size=1)
self.dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor, token_mask: torch.Tensor, min_dur: int = 1) -> Tuple[torch.Tensor, torch.Tensor]:
# x: [B, D, N], token_mask: [B, N] (bool)
B, D, N = x.shape
h = self.conv1(x) # [B, D, N]
h = F.relu(self.ln1(h.transpose(1, 2)).transpose(1, 2))
h = self.dropout(h)
h = self.conv2(h)
h = F.relu(self.ln2(h.transpose(1, 2)).transpose(1, 2))
h = self.dropout(h)
raw = self.proj(h).squeeze(1) # [B, N]
log_dur = raw * token_mask.float() #
# positive duration: softplus + min_dur
#dur_float = F.softplus(raw) + (min_dur + 1e-8)
# zero-out padding tokens
dur_float = raw * token_mask.float()
# discrete durations (ints) for upsampling and shift
dur_int = torch.clamp(torch.round(dur_float), min=min_dur).long()
dur_int = dur_int * token_mask.long()
return dur_int, dur_float
def gaussian_upsample(
token_feats: torch.Tensor, # [B, D, N]
durations: torch.Tensor, # [B, N] (int or float) - can be fractional
T: torch.Tensor, # [B] frames per sample
sigma_sq: float = 10.0,
token_mask: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Numerically-stable gaussian upsampling.
Uses a softmax over negative squared distances (log-sum-exp stability).
"""
B, D, N = token_feats.shape
device = token_feats.device
# --- MODIFIED: Handle None token_mask by calling length_to_mask_second ---
if token_mask is None:
# Assume full length N is valid for all samples if mask is missing
full_lengths = torch.full((B,), N, device=device, dtype=torch.long)
token_mask = length_to_mask_second(full_lengths, max_len=N)
# -------------------------------------------------------------------------
T_list = [int(T[b].item()) for b in range(B)]
max_T = max(T_list) if len(T_list) > 0 else 0
out_list = []
weights_list = []
for b in range(B):
Nb = int(token_mask[b].sum().item())
Tb = T_list[b]
if Nb <= 0 or Tb <= 0:
out_list.append(torch.zeros(D, max_T, device=device, dtype=token_feats.dtype))
weights_list.append(torch.zeros(N, max_T, device=device, dtype=token_feats.dtype))
continue
E = token_feats[b, :, :Nb] # [D, Nb]
d = durations[b, :Nb].float() # [Nb]
# avoid exact zeros in denom but do it smoothly with eps
d_sum = d.sum().clamp(min=1e-8)
scale = Tb / d_sum
l = d * scale # [Nb] (fractional frames)
# centers
c = torch.cumsum(l, dim=0) - 0.5 * l # [Nb]
t = torch.arange(Tb, device=device).float().unsqueeze(0) # [1, Tb]
# squared distances
dist2 = (t - c.unsqueeze(1)) ** 2 # [Nb, Tb]
# stable softmax along token axis (dim=0)
logits = (-dist2 / float(sigma_sq)).to(torch.float32) # compute in fp32
w = F.softmax(logits, dim=0).to(dist2.dtype) # [Nb, Tb]
# upsample
F_bt = E @ w # [D, Tb]
# pad to max_T
if Tb < max_T:
F_pad = torch.zeros(D, max_T, device=device, dtype=F_bt.dtype)
F_pad[:, :Tb] = F_bt
F_bt = F_pad
out_list.append(F_bt)
# pack back to [N, max_T] with padding zeros
w_full = torch.zeros(N, max_T, device=device, dtype=w.dtype)
w_full[:Nb, :Tb] = w
weights_list.append(w_full)
frame_feats = torch.stack(out_list, dim=0) # [B, D, max_T]
weights = torch.stack(weights_list, dim=0) # [B, N, max_T]
return frame_feats, weights
def alternating_shift(dur_int: torch.Tensor, alpha: int = 2, token_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
"""
Apply alternating +/- alpha shift to durations while trying to keep sum constant.
dur_int: [B, N] ints
"""
B, N = dur_int.shape
device = dur_int.device
idx = torch.arange(N, device=device)
signs = torch.where((idx % 2) == 0, 1, -1).view(1, N).expand(B, -1) # [B, N]
if token_mask is not None:
signs = signs * token_mask.long()
ds = dur_int.to(torch.long) + alpha * signs
# clip to >=1 on valid tokens
if token_mask is not None:
ds = torch.where(token_mask, torch.clamp(ds, min=1), torch.zeros_like(ds))
else:
ds = torch.clamp(ds, min=1)
# adjust last valid token to keep sum equal (compensate any clipping drift)
if token_mask is not None:
valid_counts = token_mask.sum(dim=1) # [B]
else:
valid_counts = torch.full((B,), N, device=device, dtype=torch.long)
for b in range(B):
Nb = int(valid_counts[b].item())
if Nb <= 0:
continue
diff = int(ds[b, :Nb].sum().item() - dur_int[b, :Nb].sum().item())
if diff != 0:
# subtract diff from the last token (keep >=1)
j = Nb - 1
ds[b, j] = max(int(ds[b, j].item() - diff), 1)
return ds
class MelPerFrame:
"""
Simple per-frame mel extractor returning log-mel for a mono waveform in [-1,1].
Use the same parameters you used to build ground-truth mels if you can.
"""
def __init__(self, sr: int = 44100, n_fft: int = 4096, hop: int = 600, win: int = 2400, n_mels: int = 128):
import torchaudio
self.mel = torchaudio.transforms.MelSpectrogram(
sample_rate=sr, n_fft=n_fft, hop_length=hop, win_length=win, n_mels=n_mels
).to('cuda')
self.eps = 1e-5
@torch.no_grad()
def __call__(self, wav: torch.Tensor) -> torch.Tensor:
# wav: [B, T] float
m = self.mel(wav.float()) # [B, n_mels, T_m]
m = torch.log(m + self.eps)
return m
def per_phoneme_losses_from_frame_losses(
frame_loss: torch.Tensor, # [B, T]
weights: torch.Tensor, # [B, N, T] (may be padded to max_T)
token_mask: torch.Tensor, # [B, N] (bool)
) -> torch.Tensor:
"""
Average frame losses over frames weighted by token responsibilities.
Handles potential mismatch between weights' time dim and frame_loss time dim
by using only the overlapping prefix (min time length).
Returns [B, N], zeros where token has no coverage or token is padding.
"""
B, N, Tw = weights.shape
if frame_loss.dim() != 2 or frame_loss.shape[0] != B:
raise ValueError("frame_loss must be [B, T] with same B as weights")
Tf = frame_loss.shape[1]
T_use = min(Tw, Tf)
if T_use == 0:
return torch.zeros(B, N, device=weights.device, dtype=weights.dtype)
# Trim to overlapping frames
w = weights[..., :T_use] # [B, N, T_use]
fl = frame_loss[..., :T_use] # [B, T_use]
# weighted average per token
num = (w * fl.unsqueeze(1)).sum(dim=2) # [B, N]
den_raw = w.sum(dim=2) # [B, N]
# safe division: only divide where den_raw > 0
den = den_raw + 1e-8
L = num / den
# zero tokens that had no coverage in the overlapped window
L = torch.where(den_raw > 1e-8, L, torch.zeros_like(L))
# zero out padding tokens
L = L * token_mask.float()
return L
class ReinforceAligner(nn.Module):
"""
End-to-end encoder + RL duration aligner with Gaussian upsampling.
Now: uses provided ground-truth durations to compute per-phoneme rewards.
"""
def __init__(
self,
phoneme_dim: int, # input embedding dim if you pass x_emb; or vocab size if you pass token ids
out_channels: int = 512, # encoder output dim (must match your decoder's ASR_down input)
sigma_sq: float = 10.0,
min_dur: int = 1,
num_mrf_blocks: int = 4,
use_input_is_embeds: bool = True, # True if you pass embeddings [B, N, D]; False if you pass token ids
vocab_size: int = 178, # used only if use_input_is_embeds=False (i.e., you pass ids)
dropout: float = 0.1,
):
super().__init__()
self.use_input_is_embeds = use_input_is_embeds
if not use_input_is_embeds:
self.embed = nn.Embedding(vocab_size, phoneme_dim)
self.encoder = MRFEncoder1D(in_channels=phoneme_dim, out_channels=out_channels,
num_blocks=num_mrf_blocks, dropout=dropout)
self.duration_predictor = DurationPredictor(channels=out_channels, dropout=dropout)
self.sigma_sq = sigma_sq
self.min_dur = min_dur
def inference(
self,
x: torch.Tensor,
token_lengths: torch.Tensor,
):
"""
Inference-time upsampling.
Args:
x: [B, D, N] if using embeddings; otherwise token ids (handled by self.use_input_is_embeds).
token_lengths: [B] token counts.
frame_lengths: optional [B] desired output frames length (e.g., mel length). If None,
we use rounded sum of predicted durations.
use_d_gt: optional boolean mask [B] to force using gt durations if you pass them (not used here).
with_weights: return weights matrix if True.
Returns:
dict with keys:
- "enc": encoder features [B, D, N]
- "dur_int": integer durations from predictor [B, N]
- "dur_log": raw predictor output (log space) [B, N]
- "dur": unlogged durations (frames) [B, N] = exp(dur_log)-1
- "frame_feats": upsampled frame features [B, D, T_max]
- "weights": responsibilities [B, N, T_max] (only if with_weights True)
- "token_mask": [B, N] bool
"""
device = x.device
B = x.size(0)
# embeddings if needed
if self.use_input_is_embeds:
x_emb = x # [B, D, N]
else:
x_emb = self.embed(x).transpose(1, 2) # [B, D, N]
token_mask = length_to_mask_second(token_lengths, max_len=x_emb.size(2)) # [B, N]
maskf = token_mask.float()
# encoder
enc = self.encoder(x_emb, token_mask=token_mask) # [B, D, N]
# durations from predictor (raw/log and integer)
dur_int, dur_log = self.duration_predictor(enc, token_mask, min_dur=self.min_dur) # [B, N], [B, N]
dur_log = dur_log * maskf
# convert to linear durations used in forward()
dur_unlogged = (torch.exp(dur_log) - 1.0) * maskf # [B, N]
T_pred = dur_unlogged.sum(dim=1).clamp(min=1.0).round().long().to(device) # [B]
# optionally, if you want to allow forcing GT durations at inference, you can pass them
# via use_d_gt and gt_durations; this example doesn't handle gt_durations here.
frame_feats, weights = gaussian_upsample(
enc,
dur_unlogged,
T_pred,
sigma_sq=self.sigma_sq,
token_mask=token_mask,
)
return frame_feats ,dur_unlogged , T_pred
def forward(
self,
x: torch.Tensor, # [B, D, N] if use_input_is_embeds else [B, N] token ids
token_lengths: torch.Tensor, # [B] (number of tokens per sample)
frame_lengths: torch.Tensor, # [B] (number of frames per sample; e.g., mel length)
gt_durations: Optional[torch.Tensor] = None, # [B, N] ground-truth durations per token (frames)
use_d_gt: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor, Dict]:
'''
if gt_durations is None:
raise ValueError("gt_durations must be provided when using GT-duration rewards.")
device = x.device
B = x.size(0)
# Prepare embeddings & token mask (compute token_mask after embedding so shapes are stable)
if self.use_input_is_embeds:
x_emb = x # [B, D, N]
else:
x_emb = self.embed(x).transpose(1, 2) # [B, D, N]
token_mask = length_to_mask_second(token_lengths, max_len=x_emb.size(2)) # [B, N]
# encoder
enc = self.encoder(x_emb, token_mask=token_mask) # [B, D, N]
if use_d_gt is None:
# Default to False for all if not provided
use_d_gt = torch.zeros(x.size(0), dtype=torch.bool, device=x.device)
dur_int, dur_float = self.duration_predictor(enc, token_mask, min_dur=self.min_dur) # [B, N]
maskf = token_mask.float()
#dur_float = torch.clamp(dur_float, min=float(self.min_dur))
dur_float = dur_float * maskf # [B, N]
eps = 1e-6
maskf = token_mask.float() # [B, N]
d = dur_float * maskf # [B, N]
den = d.sum(dim=1, keepdim=True) # [B, 1]
# smooth proportions (stable even if den is tiny)
p = d / (den + eps) # [B, N]
Tb = frame_lengths.float().unsqueeze(1) # [B, 1]
l_keep = p * Tb # [B, N] scaled durations that sum ~ Tb
l_keep = l_keep * maskf
mse_raw = (dur_float - (gt_durations.float()* maskf)) ** 2
#mse_raw = (dur_float - (gt_durations.float() * maskf)).abs()
mse_raw = mse_raw * maskf
denom_mse_per = maskf.sum(dim=1)
L_mse_per = mse_raw.sum(dim=1) / denom_mse_per
L_mse = L_mse_per.mean()
dur_sum = (dur_float * maskf).sum(dim=1) # [B]
L_len = ((frame_lengths.float() - dur_sum) ** 2).mean()
upsampling_durations = torch.empty_like(dur_float)
for i in range(B):
if use_d_gt[i]:
upsampling_durations[i] = gt_durations[i].float()
else:
upsampling_durations[i] = dur_float[i]
frame_feats_keep, weights_keep = gaussian_upsample(
enc,
upsampling_durations,
frame_lengths,
sigma_sq=self.sigma_sq,
token_mask=token_mask,
)
info = {
"L_mse": L_mse,
"L_len": L_len,
"dur_float": dur_float,
}
return l_keep, info , frame_feats_keep , weights_keep
'''
if gt_durations is None:
raise ValueError("gt_durations must be provided when using GT-duration rewards.")
device = x.device
B = x.size(0)
# Prepare embeddings & token mask (compute token_mask after embedding so shapes are stable)
if self.use_input_is_embeds:
x_emb = x # [B, D, N]
else:
x_emb = self.embed(x).transpose(1, 2) # [B, D, N]
token_mask = length_to_mask_second(token_lengths, max_len=x_emb.size(2)) # [B, N]
# encoder
enc = self.encoder(x_emb, token_mask=token_mask) # [B, D, N]
if use_d_gt is None:
# Default to False for all if not provided
use_d_gt = torch.zeros(x.size(0), dtype=torch.bool, device=x.device)
dur_int, dur_float = self.duration_predictor(enc, token_mask, min_dur=self.min_dur) # [B, N]
maskf = token_mask.float()
#dur_float = torch.clamp(dur_float, min=float(self.min_dur))
dur_float = dur_float * maskf # [B, N]
dur_float_unlogged = (torch.exp(dur_float) - 1) * maskf
log_gt = torch.log(gt_durations.float() + 1) * maskf
mse_raw = (dur_float - log_gt) ** 2
#mse_raw = F.huber_loss(dur_float, log_gt, delta=0.4, reduction="none")
loss_lin_l1 = torch.abs(dur_float_unlogged - gt_durations)
pause_weight = 1.0 + (gt_durations / 33.0) * 2.0
# Final Combination
# We use a higher lambda (0.1) because your max_dur is small (33),
# so gradients won't explode like they would if max_dur was 1000.
loss_per_token = (mse_raw + (0.2 * loss_lin_l1 * pause_weight)) * maskf
mse_raw = mse_raw * maskf
#denom_mse = maskf.sum().clamp(min=1.0)
denom_mse_per = maskf.sum(dim=1).clamp(min=1.0) # [B]
# per-sample MSE: sum over tokens / valid_token_count
L_mse_per = loss_per_token.sum(dim=1) / denom_mse_per # [B]
L_mse = L_mse_per.mean()
#print(denom_mse_per)
#L_mse = mse_raw.sum() / denom_mse
dur_sum = (dur_float_unlogged * maskf).sum(dim=1) # [B]
L_len = ((frame_lengths.float() - dur_sum) ** 2).mean()
upsampling_durations = torch.empty_like(dur_float_unlogged)
for i in range(B):
if use_d_gt[i]:
upsampling_durations[i] = gt_durations[i].float()
else:
upsampling_durations[i] = dur_float_unlogged[i]
frame_feats_keep, weights_keep = gaussian_upsample(
enc,
upsampling_durations,
frame_lengths,
sigma_sq=self.sigma_sq,
token_mask=token_mask,
)
info = {
"L_mse": L_mse,
"L_len": L_len,
"dur_float": dur_float,
}
eps = 1e-6
maskf = token_mask.float() # [B, N]
d = dur_float_unlogged * maskf # [B, N]
den = d.sum(dim=1, keepdim=True) # [B, 1]
# smooth proportions (stable even if den is tiny)
p = d / (den + eps) # [B, N]
Tb = frame_lengths.float().unsqueeze(1) # [B, 1]
l_keep = p * Tb # [B, N] scaled durations that sum ~ Tb
l_keep = l_keep * maskf
return l_keep, info , frame_feats_keep , weights_keep