Download models.py from FashionFlora/SFlowTTS: direct link, hf CLI and curl.
- Browser
- Download file 6.1 kB
-
https://huggingface.co/FashionFlora/SFlowTTS/resolve/main/models.py
- Command line
-
hf download hf://FashionFlora/SFlowTTS/models.py
-
curl -L -o models.py https://huggingface.co/FashionFlora/SFlowTTS/resolve/main/models.py
6.1 kB
| # coding:utf-8 | |
| import os | |
| import os.path as osp | |
| import copy | |
| import math | |
| import random | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from torch.nn.utils import spectral_norm | |
| from torch.nn.utils.parametrizations import weight_norm | |
| from collections import OrderedDict | |
| from munch import Munch | |
| import yaml | |
| import json | |
| import torchaudio | |
| from typing import Optional, Tuple, List , Sequence | |
| from Modules.text_encoder import TextEncoderTransformer | |
| from Modules.codec_decoder_hybrid_temporal import HybridTTSCodecVocoderTemporal | |
| from Modules.discriminators import MultiPeriodDiscriminator, MultiResSpecDiscriminator | |
| 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 | |
| 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()) if token_mask is not None else N | |
| 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 build_model(args): | |
| hop_length =441 | |
| fsq_levels = args.get('fsq_levels', [4] * 6) | |
| codec_strides = [2] | |
| codebook_size = np.prod(fsq_levels) | |
| codec = HybridTTSCodecVocoderSpeaker( | |
| # Mel spectrogram settings | |
| n_mels=args.n_mels, | |
| # Input dimensions | |
| text_dim=args.hidden_dim, | |
| style_dim=32, # Windowed timbre style: 32-dim vectors over ~3sec windows | |
| # Prosody latent settings | |
| prosody_latent_dim=args.prosody_latent_dim, | |
| hidden_dim=args.codec_hidden_dim, | |
| # Compression | |
| codec_strides=codec_strides, | |
| # FSQ settings | |
| codebook_size=codebook_size, | |
| fsq_levels=fsq_levels, | |
| upsample_rates=args.decoder.upsample_rates, | |
| gen_istft_n_fft=args.decoder.gen_istft_n_fft, | |
| gen_istft_hop_size=args.decoder.gen_istft_hop_size, | |
| source_upsample_rate=hop_length, | |
| language_dim=64 | |
| ) | |
| nlayers_hrm =6 | |
| H_cycles_hrm = 1 | |
| L_cycles_hrm = 3 | |
| num_heads = 8 | |
| text_encoder = TextEncoderTransformer( | |
| channels=512, | |
| language_count=7, | |
| language_hidden_dim=64, | |
| depth=6, | |
| kernel_size=5, | |
| n_symbols=args.n_token | |
| ) | |
| nets = Munch( | |
| #decoder=decoder, | |
| codec=codec, | |
| text_encoder=text_encoder, | |
| mpd = MultiPeriodDiscriminator(), | |
| msd = MultiResSpecDiscriminator(), | |
| ) | |
| return nets | |
| def load_checkpoint(model, optimizer, path, load_only_params=False, ignore_modules=[]): | |
| state = torch.load(path, map_location='cpu', weights_only=False) | |
| params = state['net'] | |
| print('loading the ckpt using the correct function.') | |
| for key in model: | |
| if key in params and key not in ignore_modules: | |
| try: | |
| model[key].load_state_dict(params[key], strict=True) | |
| except: | |
| from collections import OrderedDict | |
| state_dict = params[key] | |
| new_state_dict = OrderedDict() | |
| print(f'{key} key length: {len(model[key].state_dict().keys())}, state_dict key length: {len(state_dict.keys())}') | |
| for (k_m, v_m), (k_c, v_c) in zip(model[key].state_dict().items(), state_dict.items()): | |
| new_state_dict[k_m] = v_c | |
| model[key].load_state_dict(new_state_dict, strict=True) | |
| print('%s loaded' % key) | |
| if not load_only_params: | |
| epoch = state["epoch"] | |
| iters = state["iters"] | |
| optimizer.load_state_dict(state["optimizer"]) | |
| else: | |
| epoch = 0 | |
| iters = 0 | |
| return model, optimizer, epoch, iters |