# 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