File size: 3,129 Bytes
677a086 2a1a12e 677a086 2a1a12e 677a086 2a1a12e 677a086 2a1a12e 677a086 2a1a12e 677a086 2a1a12e 677a086 2a1a12e b4f61c2 677a086 2a1a12e 677a086 2a1a12e 677a086 2a1a12e 677a086 2a1a12e 677a086 2a1a12e 677a086 2a1a12e 677a086 2a1a12e 677a086 2a1a12e 677a086 2a1a12e 677a086 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 | import torch
import torch.nn as nn
from dataclasses import asdict
from .utils.audio import LogMelSpectrogram
from .config import ModelConfig, MelConfig
from .models.model import StableTTS
from .text import symbols
from .text import cleaned_text_to_sequence
from .text.burmese import burmese_to_ipa2
from .datas.dataset import intersperse
from .utils.audio import load_and_resample_audio
def get_vocoder(model_path, model_name='vocos'):
if model_name == 'vocos':
from .vocoders.vocos.models.model import Vocos
from .config import VocosConfig, MelConfig
vocoder = Vocos(VocosConfig(), MelConfig())
vocoder.load_state_dict(torch.load(model_path, weights_only=True, map_location='cpu'))
vocoder.eval()
else:
raise NotImplementedError(f"Unsupported vocoder: {model_name}")
return vocoder
class StableTTSAPI(nn.Module):
def __init__(self, tts_model_path, vocoder_model_path, vocoder_name='vocos'):
super().__init__()
self.mel_config = MelConfig()
self.tts_model_config = ModelConfig()
self.mel_extractor = LogMelSpectrogram(**asdict(self.mel_config))
self.tts_model = StableTTS(len(symbols), self.mel_config.n_mels, **asdict(self.tts_model_config))
if tts_model_path.endswith(".safetensors"):
from safetensors.torch import load_file
state = load_file(tts_model_path)
else:
state = torch.load(tts_model_path, map_location='cpu', weights_only=True)
self.tts_model.load_state_dict(state)
self.tts_model.eval()
self.vocoder_model = get_vocoder(vocoder_model_path, vocoder_name)
self.vocoder_model.eval()
self.g2p_mapping = {
'burmese': burmese_to_ipa2,
}
self.supported_languages = self.g2p_mapping.keys()
@torch.inference_mode()
def inference(self, text, ref_audio, language, step, temperature=1.0,
length_scale=1.0, solver=None, cfg=3.0):
device = next(self.parameters()).device
phonemizer = self.g2p_mapping.get(language)
if phonemizer is None:
raise ValueError(f"Unsupported language: {language}")
text = phonemizer(text)
text = torch.tensor(
intersperse(cleaned_text_to_sequence(text), item=0),
dtype=torch.long, device=device
).unsqueeze(0)
text_length = torch.tensor([text.size(-1)], dtype=torch.long, device=device)
ref_audio = load_and_resample_audio(ref_audio, self.mel_config.sample_rate).to(device)
ref_audio = self.mel_extractor(ref_audio)
mel_output = self.tts_model.synthesise(
text, text_length, step, temperature, ref_audio,
length_scale, solver, cfg
)['decoder_outputs']
audio_output = self.vocoder_model(mel_output)
return audio_output.cpu(), mel_output.cpu()
def get_params(self):
tts_param = sum(p.numel() for p in self.tts_model.parameters()) / 1e6
vocoder_param = sum(p.numel() for p in self.vocoder_model.parameters()) / 1e6
return tts_param, vocoder_param
|