# 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