Download Modules/durations_rl.py from FashionFlora/SFlowTTS: direct link, hf CLI and curl.
- Browser
- Download file 37.8 kB
-
https://huggingface.co/FashionFlora/SFlowTTS/resolve/main/Modules/durations_rl.py
- Command line
-
hf download hf://FashionFlora/SFlowTTS/Modules/durations_rl.py
-
curl -L -o durations_rl.py https://huggingface.co/FashionFlora/SFlowTTS/resolve/main/Modules/durations_rl.py
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 | |
| 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 | |