SFlowTTS / models.py
FashionFlora's picture
Upload full repo excluding dump_40, dump_100, precomputed_tokens, precomputed_data
fb0011a verified
Raw History Blame Contribute Delete
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