Download utils.py from FashionFlora/SFlowTTS: direct link, hf CLI and curl.
- Browser
- Download file 33.6 kB
-
https://huggingface.co/FashionFlora/SFlowTTS/resolve/main/utils.py
- Command line
-
hf download hf://FashionFlora/SFlowTTS/utils.py
-
curl -L -o utils.py https://huggingface.co/FashionFlora/SFlowTTS/resolve/main/utils.py
33.6 kB
| from monotonic_align import maximum_path | |
| from monotonic_align import mask_from_lens | |
| from monotonic_align.core import maximum_path_c | |
| import numpy as np | |
| import torch | |
| import copy | |
| from torch import nn | |
| import torch.nn.functional as F | |
| import torchaudio | |
| import librosa | |
| import matplotlib.pyplot as plt | |
| from munch import Munch | |
| from typing import List, Tuple, Optional, Union | |
| import random | |
| import re | |
| import math | |
| def maximum_path(neg_cent, mask): | |
| """ Cython optimized version. | |
| neg_cent: [b, t_t, t_s] | |
| mask: [b, t_t, t_s] | |
| """ | |
| device = neg_cent.device | |
| dtype = neg_cent.dtype | |
| neg_cent = np.ascontiguousarray(neg_cent.data.cpu().numpy().astype(np.float32)) | |
| path = np.ascontiguousarray(np.zeros(neg_cent.shape, dtype=np.int32)) | |
| t_t_max = np.ascontiguousarray(mask.sum(1)[:, 0].data.cpu().numpy().astype(np.int32)) | |
| t_s_max = np.ascontiguousarray(mask.sum(2)[:, 0].data.cpu().numpy().astype(np.int32)) | |
| maximum_path_c(path, neg_cent, t_t_max, t_s_max) | |
| return torch.from_numpy(path).to(device=device, dtype=dtype) | |
| def get_data_path_list(train_path=None, val_path=None): | |
| if train_path is None: | |
| train_path = "Data/train_list.txt" | |
| if val_path is None: | |
| val_path = "Data/val_list.txt" | |
| with open(train_path, 'r', encoding='utf-8', errors='ignore') as f: | |
| train_list = f.readlines() | |
| with open(val_path, 'r', encoding='utf-8', errors='ignore') as f: | |
| val_list = f.readlines() | |
| return train_list, val_list | |
| def length_to_mask(lengths, max_len=None): | |
| if max_len is None: | |
| max_len = lengths.max() | |
| mask = torch.arange(max_len).unsqueeze(0).expand(lengths.shape[0], -1).type_as(lengths) | |
| mask = torch.gt(mask+1, lengths.unsqueeze(1)) | |
| return mask | |
| # for norm consistency loss | |
| def log_norm1(x, mean=-4, std=4, dim=2): | |
| """ | |
| normalized log mel -> mel -> norm -> log(norm) | |
| """ | |
| x = torch.log(torch.exp(x * std + mean).norm(dim=dim)) | |
| return x | |
| def log_norm(x, mean=-4.0, std=4.0, dim=2): | |
| """ | |
| normalized log-mel x -> de-norm log-mel z -> log(sum exp(z)) over freq | |
| Returns a per-frame scalar along `dim`. | |
| """ | |
| z = x * std + mean | |
| return torch.logsumexp(z, dim=dim) / math.log(10.0) | |
| def get_image(arrs): | |
| plt.switch_backend('agg') | |
| fig = plt.figure() | |
| ax = plt.gca() | |
| ax.imshow(arrs) | |
| return fig | |
| def recursive_munch(d): | |
| if isinstance(d, dict): | |
| return Munch((k, recursive_munch(v)) for k, v in d.items()) | |
| elif isinstance(d, list): | |
| return [recursive_munch(v) for v in d] | |
| else: | |
| return d | |
| def log_print(message, logger): | |
| logger.info(message) | |
| print(message) | |
| def build_context_window_center_exists_prob( | |
| sample: Union[List[str], Tuple[str, ...], str], | |
| num_context: int = 5, | |
| drop_prob: float = 0.05, | |
| rng: Optional[np.random.RandomState] = None, | |
| return_mask: bool = True, | |
| ) -> Union[Tuple[List[str], List[bool]], List[str]]: | |
| """ | |
| Build a symmetric context window around the center with hierarchical drops: | |
| - With probability drop_prob: choose side (50/50) and neighbor (1 or 2). | |
| * If neighbor-1 chosen -> drop neighbor-1 AND neighbor-2 on that side. | |
| * If neighbor-2 chosen -> drop only neighbor-2 on that side. | |
| Center is NEVER dropped and is always marked True in the mask. | |
| Returns (window_list, orig_mask) when return_mask is True. | |
| """ | |
| if rng is None: | |
| rng = np.random | |
| # Single string -> repeat, all True | |
| if not isinstance(sample, (list, tuple)): | |
| window = [str(sample)] * num_context | |
| return (window, [True] * num_context) if return_mask else window | |
| L = len(sample) | |
| if L == 0: | |
| window = [""] * num_context | |
| return (window, [False] * num_context) if return_mask else window | |
| if num_context % 2 == 0: | |
| raise ValueError("num_context must be odd (e.g. 5)") | |
| center_idx = L // 2 # user guarantee: center exists | |
| half = num_context // 2 | |
| slot_sample_idxs = [center_idx - half + i for i in range(num_context)] | |
| # initial fill: strings or None if out-of-bounds | |
| values: List[Optional[str]] = [] | |
| src_idx: List[Optional[int]] = [] | |
| for idx in slot_sample_idxs: | |
| if 0 <= idx < L: | |
| values.append(sample[idx]) | |
| src_idx.append(idx) | |
| else: | |
| values.append(None) | |
| src_idx.append(None) | |
| center_pos = half | |
| # ensure center is set and marked original | |
| if values[center_pos] is None: | |
| values[center_pos] = sample[center_idx] | |
| src_idx[center_pos] = center_idx | |
| # Decide hierarchical drop once per sample | |
| if rng.random() < drop_prob: | |
| go_left = (rng.random() < 0.5) | |
| choose_neighbor1 = (rng.random() < 0.5) | |
| prev1_pos = center_pos - 1 | |
| prev2_pos = center_pos - 2 | |
| next1_pos = center_pos + 1 | |
| next2_pos = center_pos + 2 | |
| if go_left: | |
| if choose_neighbor1: | |
| if 0 <= prev1_pos < num_context: | |
| values[prev1_pos] = None; src_idx[prev1_pos] = None | |
| if 0 <= prev2_pos < num_context: | |
| values[prev2_pos] = None; src_idx[prev2_pos] = None | |
| else: | |
| if 0 <= prev2_pos < num_context: | |
| values[prev2_pos] = None; src_idx[prev2_pos] = None | |
| else: | |
| if choose_neighbor1: | |
| if 0 <= next1_pos < num_context: | |
| values[next1_pos] = None; src_idx[next1_pos] = None | |
| if 0 <= next2_pos < num_context: | |
| values[next2_pos] = None; src_idx[next2_pos] = None | |
| else: | |
| if 0 <= next2_pos < num_context: | |
| values[next2_pos] = None; src_idx[next2_pos] = None | |
| # Fill missing slots with fallback order: center -> next outward -> prev outward | |
| def fill_missing_slot(i: int): | |
| if values[i] not in (None, ""): | |
| return | |
| cand_positions = [center_pos] | |
| for d in range(1, num_context): | |
| p = center_pos + d | |
| if p < num_context: | |
| cand_positions.append(p) | |
| for d in range(1, num_context): | |
| p = center_pos - d | |
| if p >= 0: | |
| cand_positions.append(p) | |
| for p in cand_positions: | |
| if values[p] not in (None, ""): | |
| values[i] = values[p] | |
| return | |
| values[i] = "" | |
| for i in range(num_context): | |
| if i == center_pos: | |
| continue | |
| fill_missing_slot(i) | |
| final_window = [str(v) for v in values] | |
| if return_mask: | |
| original_mask = [si is not None for si in src_idx] | |
| # ensure center mask is True (enforce) | |
| original_mask[center_pos] = True | |
| return final_window, original_mask | |
| return final_window | |
| def mask_input_ids_random( | |
| input_ids, | |
| attention_mask, | |
| mask_token_id, | |
| p_mask=0.15, | |
| exclude_token_ids=None, | |
| seed=None, | |
| ): | |
| """ | |
| Losowo zamienia tokeny na mask_token_id z prawdopodobienstwem p_mask. | |
| Nie dotyka paddingu (attention_mask==0), ani tokenów z exclude_token_ids. | |
| input_ids: Tensor [B, L] | |
| attention_mask: Tensor [B, L] | |
| Returns: new_input_ids (cloned) | |
| """ | |
| if seed is not None: | |
| torch.manual_seed(seed) | |
| random.seed(seed) | |
| B, L = input_ids.shape | |
| new_input_ids = input_ids.clone() | |
| if exclude_token_ids is None: | |
| exclude_token_ids = set() | |
| for i in range(B): | |
| valid_mask = attention_mask[i].bool() | |
| # indices we can consider (True where mask==1 and token not excluded) | |
| cand_idx = [] | |
| for j in range(L): | |
| if not valid_mask[j]: | |
| continue | |
| tok = int(input_ids[i, j].item()) | |
| if tok in exclude_token_ids: | |
| continue | |
| cand_idx.append(j) | |
| if not cand_idx: | |
| continue | |
| # sample positions independently | |
| probs = torch.rand(len(cand_idx)) | |
| for k, pos in enumerate(cand_idx): | |
| if probs[k].item() < p_mask: | |
| new_input_ids[i, pos] = mask_token_id | |
| return new_input_ids | |
| def shuffle_segments_by_separator( | |
| input_ids, | |
| attention_mask, | |
| sep_token_id, | |
| p_shuffle=0.2, | |
| bos_token_id=None, | |
| eos_token_id=None, | |
| seed=None, | |
| ): | |
| """ | |
| Dzieli sekwencję na segmenty rozdzielone sep_token_id (w obrębie ważnej części: | |
| od pierwsnego do ostatniego tokena gdzie attention_mask==1), następnie z prawdopodobieństwem | |
| p_shuffle permutuje kolejność segmentów. BOS i EOS są zachowane na swoich pozycjach. | |
| input_ids: Tensor [B, L] | |
| Returns: new_input_ids (cloned) | |
| """ | |
| if seed is not None: | |
| random.seed(seed) | |
| B, L = input_ids.shape | |
| new_input_ids = input_ids.clone() | |
| pad_token_id = 0 # fallback, można pobrać z tokenizer.pad_token_id jeśli dostępne | |
| for i in range(B): | |
| valid_positions = (attention_mask[i] == 1).nonzero(as_tuple=True)[0].tolist() | |
| if not valid_positions: | |
| continue | |
| first_pos = valid_positions[0] | |
| last_pos = valid_positions[-1] | |
| # verify first is bos and last is eos when given (if not given, still proceed) | |
| seq = input_ids[i].tolist() | |
| # We will keep seq[first_pos] and seq[last_pos] unchanged. | |
| middle = seq[first_pos + 1 : last_pos] # list | |
| if not middle: | |
| continue | |
| # split middle into chunks where sep_token ends chunk | |
| chunks = [] | |
| current = [] | |
| for tok in middle: | |
| current.append(tok) | |
| if tok == sep_token_id: | |
| chunks.append(current) | |
| current = [] | |
| if current: | |
| chunks.append(current) | |
| # if there is only one chunk or no sep present, nothing to shuffle | |
| if len(chunks) <= 1: | |
| continue | |
| if random.random() < p_shuffle: | |
| random.shuffle(chunks) | |
| new_middle = [] | |
| for c in chunks: | |
| new_middle.extend(c) | |
| # build new sequence keeping first_pos and last_pos tokens | |
| new_seq = [] | |
| # tokens before first_pos (likely BOS and maybe specials) | |
| # keep prefix up to first_pos (inclusive) | |
| new_seq.extend(seq[: first_pos + 1]) | |
| new_seq.extend(new_middle) | |
| # then last token(s) from last_pos to end (inclusive) | |
| new_seq.extend(seq[last_pos:]) | |
| # pad / truncate to L | |
| if len(new_seq) < L: | |
| new_seq = new_seq + [pad_token_id] * (L - len(new_seq)) | |
| else: | |
| new_seq = new_seq[:L] | |
| new_input_ids[i] = torch.tensor(new_seq, dtype=input_ids.dtype, device=input_ids.device) | |
| return new_input_ids | |
| # --- UTILS: embedding-level augmentations (po gemma_mrl) --- | |
| def augment_embeddings_noise_and_dropout( | |
| embeddings, | |
| attention_mask=None, # optional: [B, T] mask aligned with embeddings if available | |
| p_noise=0.1, | |
| noise_std=0.02, | |
| p_dropout=0.05, | |
| seed=None, | |
| ): | |
| """ | |
| embeddings: Tensor [B, T, D] (np. output gemma_mrl) | |
| Dodaje gaussowski szum do losowo wybranych embeddingów i/lub ustawia | |
| niektóre embeddingi na zero (dropout). | |
| """ | |
| if seed is not None: | |
| torch.manual_seed(seed) | |
| random.seed(seed) | |
| B, T, D = embeddings.shape | |
| out = embeddings.clone() | |
| # noise | |
| if p_noise > 0: | |
| mask = torch.rand(B, T, device=embeddings.device) < p_noise | |
| if attention_mask is not None: | |
| # attention_mask: align shape [B, T] (optional) | |
| mask = mask & (attention_mask.bool()) | |
| noise = torch.randn_like(out) * noise_std | |
| out = out + noise * mask.unsqueeze(-1).to(out.dtype) | |
| # dropout (zero-out some embeddings) | |
| if p_dropout > 0: | |
| drop_mask = torch.rand(B, T, device=embeddings.device) < p_dropout | |
| if attention_mask is not None: | |
| drop_mask = drop_mask & (attention_mask.bool()) | |
| out[drop_mask] = 0.0 | |
| return out | |
| class DurationProcessor(torch.nn.Module): | |
| def __init__(self, class_count, max_dur): | |
| super(DurationProcessor, self).__init__() | |
| self.class_count = class_count | |
| self.max_dur = max_dur | |
| class_to_dur_table = torch.Tensor( | |
| [1, 2, 3, 4, 5, 6, 7, 9, 12, 15, 18, 22, 27, 32, 38, 46] | |
| ) | |
| self.register_buffer("class_to_dur_table", class_to_dur_table) | |
| dur_to_class_table = torch.Tensor( | |
| [ | |
| 0, | |
| 0, | |
| 1, | |
| 2, | |
| 3, | |
| 4, | |
| 5, | |
| 6, | |
| 7, | |
| 7, | |
| 7, | |
| 8, | |
| 8, | |
| 8, | |
| 9, | |
| 9, | |
| 9, | |
| 10, | |
| 10, | |
| 10, | |
| 11, | |
| 11, | |
| 11, | |
| 11, | |
| 11, | |
| 12, | |
| 12, | |
| 12, | |
| 12, | |
| 12, | |
| 13, | |
| 13, | |
| 13, | |
| 13, | |
| 13, | |
| 14, | |
| 14, | |
| 14, | |
| 14, | |
| 14, | |
| 14, | |
| 14, | |
| 15, | |
| 15, | |
| 15, | |
| 15, | |
| 15, | |
| 15, | |
| 15, | |
| 15, | |
| 15, | |
| ] | |
| ) | |
| self.register_buffer("dur_to_class_table", dur_to_class_table) | |
| # def class_to_dur_soft(self, class_dist): | |
| # return class_dist * self.class_to_dur_table | |
| def class_to_dur_soft(self, softdur): | |
| result = (softdur * self.class_to_dur_table).sum(dim=-1) / ( | |
| softdur.sum(dim=-1) + 1e-9 | |
| ) | |
| return result | |
| def class_to_dur_hard(self, classes): | |
| classes = classes.clamp(min=0, max=self.class_count) | |
| return self.class_to_dur_table[classes] | |
| def dur_to_class(self, durs): | |
| durs = durs.clamp(min=1, max=self.max_dur) | |
| return self.dur_to_class_table[durs.long()] | |
| def align_to_class(self, alignment): | |
| result = alignment.sum(dim=-1).clamp(min=1, max=50) | |
| result = self.dur_to_class(result) | |
| return result | |
| def prediction_to_duration(self, pred, text_length): | |
| # softdur = self.class_to_dur_soft(torch.softmax(pred, dim=-1)) | |
| # softdur = softdur.sum(dim=-1).round().clamp(min=1) | |
| # argmax = torch.argmax(pred, dim=-1).long() | |
| # argdur = self.class_to_dur_hard(argmax) | |
| confidence = torch.softmax(pred, dim=-1) | |
| softdur = self.class_to_dur_soft(confidence) | |
| dur = softdur | |
| # dur = (argdur * (argdur < 7)) + (softdur * (argdur >= 7)) | |
| # dur = dur[:text_length] | |
| return dur | |
| def _make_masks( | |
| self,B: int, S: int, T: int, text_lens: Optional[torch.Tensor], mel_lens: torch.Tensor | |
| ): | |
| device = mel_lens.device | |
| s_idx = torch.arange(S, device=device).view(1, S) | |
| t_idx = torch.arange(T, device=device).view(1, T) | |
| if text_lens is None: | |
| token_mask = torch.ones(B, S, dtype=torch.bool, device=device) | |
| else: | |
| token_mask = s_idx < text_lens.view(B, 1) | |
| frame_mask = t_idx < mel_lens.view(B, 1) | |
| return token_mask, frame_mask | |
| def align_from_softdurations( | |
| self, | |
| d_gt: torch.Tensor, # [B, S], nonnegative; trailing zeros = pads | |
| mel_lens: Optional[torch.Tensor] = None, # [B], ints (#frames per sample) | |
| tau: float = 1.0, # entropic temperature (smaller = sharper) | |
| sigma_scale: float = 0.5, # window width grows with duration | |
| sigma_floor: float = 0.5, # minimum width | |
| iters: int = 60, # Sinkhorn iterations | |
| eps: float = 1e-8, | |
| ) -> torch.Tensor: | |
| """ | |
| Reconstruct a soft alignment P from soft durations d_gt using entropic OT. | |
| Args: | |
| d_gt: [B, S] nonnegative durations (can be fractional). Zeros are treated | |
| as padded tokens if text lengths are unknown. | |
| mel_lens: [B] number of frames per sample. If None, it's inferred as | |
| round(sum(d_gt[b])) for each sample. | |
| tau: entropic regularization temperature (lower = crisper). | |
| sigma_scale, sigma_floor: shape the Gaussian cost (not the row mass). | |
| iters: Sinkhorn iterations. | |
| eps: small epsilon for numerical stability. | |
| Returns: | |
| P: [B, S, T_max] alignment. For each sample b: | |
| - sum_t P[b, i, t] == scaled d_gt[b, i] | |
| - sum_i P[b, i, t] == 1 for t < mel_lens[b], else 0 | |
| """ | |
| assert d_gt.dim() == 2 | |
| B, S = d_gt.shape | |
| device, dtype = d_gt.device, d_gt.dtype | |
| # Infer mel lengths if not provided: T_b ≈ sum of durations for that sample. | |
| if mel_lens is None: | |
| mel_lens = ( | |
| d_gt.sum(dim=1).round().clamp_min(1).to(torch.long) | |
| ) # [B] | |
| T = int(mel_lens.max().item()) | |
| # Masks | |
| token_mask = d_gt > eps # [B, S] treat zeros as padded tokens | |
| frame_idx = torch.arange(T, device=device).view(1, T) | |
| frame_mask = frame_idx < mel_lens.view(B, 1) # [B, T] | |
| token_mask3 = token_mask.unsqueeze(-1) # [B, S, 1] | |
| frame_mask3 = frame_mask.unsqueeze(1) # [B, 1, T] | |
| # Target row marginals r: scale d_gt to match mel_lens exactly per sample. | |
| d = torch.where(token_mask, d_gt, torch.zeros_like(d_gt)) # zero-out pads | |
| sum_d = d.sum(dim=1, keepdim=True).clamp_min(eps) # [B, 1] | |
| scale = mel_lens.to(dtype).view(B, 1) / sum_d # [B, 1] | |
| r = d * scale # [B, S] | |
| # Gaussian-shaped cost centered by cumulative r (monotonic prior). | |
| end = torch.cumsum(r, dim=1) # [B, S] | |
| centers = (end - 0.5 * r).unsqueeze(-1) # [B, S, 1] | |
| sigma = sigma_scale * r.unsqueeze(-1) + sigma_floor # [B, S, 1] | |
| t = torch.arange(T, device=device, dtype=dtype).view(1, 1, T) | |
| x = t - centers | |
| cost = 0.5 * (x / (sigma + eps)) ** 2 # [B, S, T] | |
| big = 1e6 | |
| cost = cost + (~token_mask3) * big + (~frame_mask3) * big | |
| logK = -cost / tau # kernel in log-space | |
| # Log-marginals | |
| log_r = torch.full((B, S), -float("inf"), device=device, dtype=dtype) | |
| valid_r = r > 0 | |
| log_r[valid_r] = torch.log(r[valid_r]) | |
| log_c = torch.full((B, T), -float("inf"), device=device, dtype=dtype) | |
| log_c[frame_mask] = 0.0 # log(1) for valid frames | |
| # Sinkhorn iterations (log-domain) | |
| log_u = torch.zeros((B, S), device=device, dtype=dtype) | |
| log_v = torch.zeros((B, T), device=device, dtype=dtype) | |
| for _ in range(iters): | |
| logMv = torch.logsumexp(logK + log_v.unsqueeze(1), dim=2) # [B, S] | |
| log_u = log_r - logMv | |
| logKu = torch.logsumexp(logK + log_u.unsqueeze(2), dim=1) # [B, T] | |
| log_v = log_c - logKu | |
| logP = log_u.unsqueeze(2) + logK + log_v.unsqueeze(1) # [B, S, T] | |
| P = torch.exp(logP) * token_mask3 * frame_mask3 | |
| return P | |
| def duration_to_alignment_sinkhorn_batched( | |
| self, | |
| durations: torch.Tensor, # [B, S], nonnegative | |
| mel_lens: torch.Tensor, # [B], ints (#frames per sample) | |
| text_lens: Optional[torch.Tensor] = None, # [B], ints (valid tokens) | |
| tau: float = 1.0, # entropic temperature (smaller = sharper) | |
| sigma_scale: float = 0.5, # width grows with duration | |
| sigma_floor: float = 0.5, # minimum width | |
| iters: int = 60, # Sinkhorn iterations | |
| eps: float = 1e-8, | |
| ) -> torch.Tensor: | |
| """ | |
| Returns P in [B, S, T_max] s.t. | |
| - sum_t P[b, i, t] == r[b, i] (scaled durations) | |
| - sum_i P[b, i, t] == 1 for t < mel_lens[b], else 0 | |
| Fully differentiable w.r.t. durations. | |
| """ | |
| assert durations.dim() == 2 | |
| B, S = durations.shape | |
| device = durations.device | |
| dtype = durations.dtype | |
| T = int(mel_lens.max().item()) | |
| # Masks | |
| token_mask, frame_mask = self._make_masks(B, S, T, text_lens, mel_lens) | |
| token_mask3 = token_mask.unsqueeze(-1) # [B, S, 1] | |
| frame_mask3 = frame_mask.unsqueeze(1) # [B, 1, T] | |
| # Build row marginals r: scale durations to sum to mel_lens per sample. | |
| d = durations.clamp_min(eps) | |
| if text_lens is not None: | |
| d = d * token_mask.float() | |
| sum_d = d.sum(dim=1, keepdim=True).clamp_min(eps) # [B, 1] | |
| scale = mel_lens.to(dtype).view(B, 1) / sum_d # [B, 1] | |
| r = d * scale # [B, S], sums to mel_len per sample on valid tokens | |
| # Zero-out padded tokens exactly in r | |
| if text_lens is not None: | |
| r = r * token_mask.float() | |
| # Gaussian-shaped cost around differentiable centers from r | |
| end = torch.cumsum(r, dim=1) # [B, S] | |
| centers = (end - 0.5 * r).unsqueeze(-1) # [B, S, 1] | |
| sigma = sigma_scale * r.unsqueeze(-1) + sigma_floor # [B, S, 1] | |
| t = torch.arange(T, device=device, dtype=dtype).view(1, 1, T) | |
| x = t - centers | |
| # Cost ~ 0.5 * (x/sigma)^2 (no absolute scale needed; tau controls sharpness) | |
| cost = 0.5 * (x / (sigma + eps)) ** 2 # [B, S, T] | |
| big = 1e6 | |
| if text_lens is not None: | |
| cost = cost + (~token_mask3) * big | |
| cost = cost + (~frame_mask3) * big | |
| logK = -cost / tau # [B, S, T] | |
| # Log marginals | |
| log_r = torch.full((B, S), -float("inf"), device=device, dtype=dtype) | |
| valid_r = r > 0 | |
| log_r[valid_r] = torch.log(r[valid_r]) | |
| log_c = torch.full((B, T), -float("inf"), device=device, dtype=dtype) | |
| log_c[frame_mask] = 0.0 # log(1) for valid frames | |
| # Sinkhorn iterations in log-domain | |
| log_u = torch.zeros((B, S), device=device, dtype=dtype) | |
| log_v = torch.zeros((B, T), device=device, dtype=dtype) | |
| for _ in range(iters): | |
| # log_u = log_r - logsumexp_j [logK + log_v_j] | |
| logMv = torch.logsumexp(logK + log_v.unsqueeze(1), dim=2) # [B, S] | |
| log_u = log_r - logMv | |
| # log_v = log_c - logsumexp_i [logK + log_u_i] | |
| logKu = torch.logsumexp(logK + log_u.unsqueeze(2), dim=1) # [B, T] | |
| log_v = log_c - logKu | |
| # Transport plan | |
| logP = log_u.unsqueeze(2) + logK + log_v.unsqueeze(1) # [B, S, T] | |
| P = torch.exp(logP) | |
| # Mask out pads for cleanliness (mass is already ~0 there) | |
| P = P * token_mask3 * frame_mask3 | |
| return P | |
| def duration_to_alignment_soft( | |
| self, | |
| duration: torch.Tensor, | |
| width_factor: float =1.0, | |
| min_width: float = 1.0, | |
| mask_margin: int = 2, | |
| temperature: float = 0.5, | |
| ) -> torch.Tensor: | |
| """ | |
| Convert durations -> attention matrix. | |
| duration: (T,) or (B, T) of positive frame counts (floats or ints). | |
| Returns: (T, A) if input was (T,), else (B, T, A). | |
| T = text length, A = audio frames (sum of durations). | |
| width_factor controls the per-token sigma ~ duration * width_factor. | |
| min_width prevents sigma from becoming too small. | |
| temperature < 1 sharpens the softmax (use 1.0 for no change). | |
| mask_margin extends the token window by this many frames. | |
| """ | |
| single_example = False | |
| if duration.dim() == 1: | |
| duration = duration.unsqueeze(0) # -> (1, T) | |
| single_example = True | |
| device = duration.device | |
| duration = duration.to(dtype=torch.float32, device=device) | |
| # total audio length (A) | |
| total_dur = int(duration.sum(dim=1).round().max().item()) | |
| if total_dur <= 0: | |
| total_dur = 1 | |
| # per-token cumulative bounds and centers | |
| upper = torch.cumsum(duration, dim=1) # (B, T) | |
| lower = upper - duration | |
| mean = (lower + upper) / 2.0 # (B, T) | |
| mean = mean.unsqueeze(2) # (B, T, 1) | |
| # frame indices (1, 1, A) | |
| seq = torch.arange(total_dur, device=device, dtype=duration.dtype) | |
| seq = seq.unsqueeze(0).unsqueeze(0) # (1, 1, A) | |
| # distance from center (B, T, A) | |
| x = seq - mean | |
| # per-token sigma (B, T, 1) | |
| sigma = (duration.unsqueeze(2) * width_factor) + min_width | |
| sigma = sigma.clamp(min=1e-6) | |
| # logits = -0.5 * (x / sigma)**2 (Gaussian log-likelihood up to const) | |
| logits = -0.5 * (x / sigma) ** 2 | |
| # mask out-of-bound frames by setting logits to a very large negative value | |
| lower_mask = seq >= (lower - mask_margin).unsqueeze(2) | |
| upper_mask = seq <= (upper + mask_margin).unsqueeze(2) | |
| mask = (lower_mask & upper_mask).to(dtype=torch.bool) | |
| logits = logits.masked_fill(~mask, -1e9) | |
| # optional sharpening via temperature: smaller -> sharper | |
| if temperature <= 0: | |
| raise ValueError("temperature must be > 0") | |
| logits = logits / float(temperature) | |
| # softmax over text positions (dim=1) so each audio frame has a distribution over tokens | |
| alignment = torch.softmax(logits, dim=1) | |
| if single_example: | |
| return alignment[0] # (T, A) | |
| return alignment # (B, T, A) | |
| def duration_to_alignment(self, duration: torch.Tensor) -> torch.Tensor: | |
| indices = torch.repeat_interleave( | |
| torch.arange(duration.shape[0], device=duration.device), | |
| duration.to(torch.int), | |
| ) | |
| result = torch.zeros( | |
| (duration.shape[0], indices.shape[0]), device=duration.device | |
| ) | |
| result[indices, torch.arange(indices.shape[0])] = 1 | |
| return result | |
| def forward(self, pred, text_length): | |
| duration = self.prediction_to_duration(pred, text_length) | |
| alignment = self.duration_to_alignment_soft(duration) | |
| return 0 | |
| def torch_empty_cache(device): | |
| if device == "cuda": | |
| torch.cuda.synchronize() | |
| torch.cuda.empty_cache() | |
| elif device == "mps": | |
| torch.mps.synchronize() | |
| torch.mps.empty_cache() | |
| elif device == "cpu": | |
| # torch.cpu.synchronize() | |
| # torch.cpu.empty_cache() | |
| pass | |
| else: | |
| exit(f"Unknown device {device}. Could not empty cache.") | |
| SENTENCE_SPLIT_RE = re.compile(r'(?<=[\.\?\!])\s+|\n+') | |
| def split_sentences(text: Optional[str]) -> List[str]: | |
| """Dzieli tekst na zdania; jeśli brak typowych końców zdań, | |
| traktuje cały tekst jako jedno zdanie. Zwraca listę zdań.""" | |
| if text is None: | |
| return [] | |
| text = text.strip() | |
| if not text: | |
| return [] | |
| parts = [s.strip() for s in SENTENCE_SPLIT_RE.split(text) if s.strip()] | |
| if not parts: | |
| return [text] | |
| return parts | |
| def select_context(prev2_txt: str, prev1_txt: str, curr_txt: str, | |
| next_txt: str) -> List[str]: | |
| """Zwraca listę zdań do użycia jako embeddingi zgodnie z regułami: | |
| - jeśli curr >=4 -> wszystkie zdania z curr | |
| - jeśli curr ==3 -> prev2=curr[0], prev1=curr[1], curr=curr[2], next=first(next) | |
| - jeśli curr ==2 -> prev2=last(prev1) (fallback z prev2), prev1=curr[0], | |
| curr=curr[1], next=last(next) | |
| - jeśli curr ==1 -> prev1/prev2 zgodnie z regułami opisanymi przez Ciebie, | |
| next=first(next) | |
| - jeśli curr ==0 -> fallback: ostatnie 4 zdania z (prev2+prev1+next) lub | |
| puste stringi | |
| """ | |
| p2 = split_sentences(prev2_txt) | |
| p1 = split_sentences(prev1_txt) | |
| c = split_sentences(curr_txt) | |
| n = split_sentences(next_txt) | |
| # case curr >= 4: wszystkie zdania z curr jako embeddingi | |
| if len(c) >= 4: | |
| return c | |
| if len(c) == 3: | |
| prev2 = p1[-1] if p1 else "" | |
| prev1 = c[0] | |
| curr = c[1] | |
| #next_s = n[0] if n else "" | |
| next_s = c[2] | |
| return [prev2, prev1, curr, next_s] | |
| if len(c) == 2: | |
| prev1_from_curr0 = c[0] | |
| curr_sent = c[1] | |
| # prev2: ostatnie zdanie z prev1, fallback do ostatniego z prev2 | |
| if len(p1) >= 1: | |
| prev2_from_p1 = p1[-1] | |
| elif len(p2) >= 1: | |
| prev2_from_p1 = p2[-1] | |
| else: | |
| prev2_from_p1 = "" | |
| next_s = n[-1] if n else "" | |
| return [prev2_from_p1, prev1_from_curr0, curr_sent, next_s] | |
| if len(c) == 1: | |
| curr_sent = c[0] | |
| # prev1 / prev2 zależnie od długości prev1 | |
| if len(p1) >= 2: | |
| prev1_s = p1[-1] | |
| prev2_s = p1[-2] | |
| elif len(p1) == 1: | |
| prev1_s = p1[0] | |
| prev2_s = p2[-1] if p2 else "" | |
| else: | |
| # brak prev1, spróbuj wypełnić z prev2 | |
| if len(p2) >= 2: | |
| prev1_s = p2[-1] | |
| prev2_s = p2[-2] | |
| elif len(p2) == 1: | |
| prev1_s = p2[-1] | |
| prev2_s = "" | |
| else: | |
| prev1_s = "" | |
| prev2_s = "" | |
| next_s = n[0] if n else "" | |
| return [prev2_s, prev1_s, curr_sent, next_s] | |
| # fallback (curr == 0): we bierzemy ostatnie 4 zdania z prev2+prev1+next | |
| all_sents = p2 + p1 + n | |
| if not all_sents: | |
| return ["", "", "", ""] | |
| res = all_sents[-4:] | |
| # jeśli mniej niż 4 - dopadamy pustymi stringami z lewej | |
| while len(res) < 4: | |
| res.insert(0, "") | |
| return res | |
| def build_s2s_attn_from_predictor_no_max( | |
| d_preds: List[torch.Tensor], | |
| input_lengths: torch.Tensor, | |
| l_min: int = 1, | |
| l_max: Optional[int] = 8000, | |
| window_radius_tokens: int = 2, # focus: 1–3 is typical | |
| token_temp: float = 0.6, # focus: <1 sharper, >1 softer | |
| kappa: float = 0.8, # how much to use predictor modulation | |
| add_uniform_eps: float = 1e-4, # avoids dead columns | |
| ) -> Tuple[torch.Tensor, List[int]]: | |
| """ | |
| Builds coarse-but-focused attention: | |
| - frame-to-token path j_hat from cumulative durations (coarse across frames) | |
| - restrict tokens to a small window around j_hat (focused across tokens) | |
| - sharpen inside the window with a temperature | |
| - optionally modulate by predictor signal resampled to frame grid | |
| Returns: | |
| s2s_attn: [B, T_max, L_batch_max], column-stochastic | |
| output_lengths: list[int] | |
| """ | |
| device = input_lengths.device | |
| B = int(input_lengths.numel()) | |
| T_max = int(input_lengths.max().item()) | |
| output_lengths: List[int] = [] | |
| cache = [] | |
| # Pass 1: durations + predicted output length | |
| for pred, t_len in zip(d_preds, input_lengths): | |
| T = int(t_len.item()) | |
| logits = pred[:T].to(device, dtype=torch.float32) # [T, L_pred] | |
| probs = torch.sigmoid(logits) # [T, L_pred] | |
| dur = probs.sum(dim=-1) + 1e-6 # [T] | |
| l_float = dur.sum() | |
| if l_max is None: | |
| l = int(torch.clamp(l_float.round(), min=l_min).item()) | |
| else: | |
| l = int(torch.clamp(l_float.round(), min=l_min, max=l_max).item()) | |
| # Scale so sum(dur) == l (approximately exact) | |
| scale = float(l) / (l_float.item() + 1e-8) | |
| dur = dur * scale | |
| output_lengths.append(l) | |
| cache.append((logits, dur, T, l)) | |
| L_batch_max = max(output_lengths) if output_lengths else l_min | |
| s2s = torch.zeros(B, T_max, L_batch_max, device=device, dtype=torch.float32) | |
| # Pass 2: build windowed attention around path j_hat | |
| for b, (logits, dur, T, l) in enumerate(cache): | |
| # frame midpoints and cumulative duration to get token index per frame | |
| cdf = torch.cumsum(dur, dim=0) # [T] | |
| frames = torch.arange(l, device=device, dtype=cdf.dtype) + 0.5 # [l] | |
| j_hat = torch.searchsorted(cdf, frames) # [l], long | |
| j_hat = j_hat.clamp(min=0, max=T - 1) | |
| # Token indices grid | |
| t_idx = torch.arange(T, device=device).unsqueeze(1).to(cdf.dtype) # [T,1] | |
| j_hat_f = j_hat.unsqueeze(0).to(cdf.dtype) # [1,l] | |
| # Window mask around current token index | |
| R = int(window_radius_tokens) | |
| low = (j_hat - R).clamp(min=0) | |
| high = (j_hat + R).clamp(max=T - 1) | |
| # allowed[t,f] = low[f] <= t <= high[f] | |
| allowed = (t_idx >= low.unsqueeze(0)) & (t_idx <= high.unsqueeze(0)) # [T,l] | |
| # Base logits: distance in token index from j_hat (focus within window) | |
| # Use a small Gaussian over token index; sigma ~ R / 1.5 | |
| sigma_tokens = max(1.0, R / 1.5) | |
| attn_logits = -0.5 * ((t_idx - j_hat_f) / sigma_tokens) ** 2 # [T,l] | |
| # Optional modulation by predictor (resampled to l frames) | |
| if kappa > 0.0: | |
| pred_res = F.interpolate( | |
| logits.unsqueeze(0), size=l, mode="linear", align_corners=True | |
| ).squeeze(0) # [T,l] | |
| attn_logits = attn_logits + kappa * torch.log( | |
| torch.sigmoid(pred_res) + 1e-6 | |
| ) | |
| # Mask outside the window | |
| attn_logits = attn_logits.masked_fill(~allowed, -1e4) | |
| # Sharpen within window with temperature | |
| attn = F.softmax(attn_logits / max(1e-4, token_temp), dim=0) # [T,l] | |
| # Tiny uniform prior to avoid degenerate columns | |
| if add_uniform_eps > 0: | |
| attn = attn + add_uniform_eps | |
| attn = attn / (attn.sum(dim=0, keepdim=True) + 1e-8) | |
| s2s[b, :T, :l] = attn.to(s2s.dtype) | |
| return s2s, output_lengths |