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