MyanmarTTS / api.py
freococo's picture
Sync api.py with clean PyPI version
2a1a12e verified
Raw History Blame Contribute Delete
3.13 kB
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