Aloento/9Nine-PITS
1
1import argparse2 3import gradio as gr4import torch5 6import commons7import utils8from models import SynthesizerTrn9from text import cleaned_text_to_sequence10from text.cleaners import clean_text11from text.symbols import symbols12 13 14# we use Kyubyong/g2p for demo instead of our internal g2p15# https://github.com/Kyubyong/g2p16def get_text(text, hps):17 cleaned_text, lang = clean_text(text)18 text_norm = cleaned_text_to_sequence(cleaned_text)19 if hps.data.add_blank:20 text_norm, lang = commons.intersperse_with_language_id(text_norm, lang, 0)21 text_norm = torch.LongTensor(text_norm)22 lang = torch.LongTensor(lang)23 return text_norm, lang, cleaned_text24 25 26class GradioApp:27 28 def __init__(self, args):29 self.hps = utils.get_hparams_from_file(args.config)30 self.device = "cpu"31 32 self.net_g = SynthesizerTrn(33 len(symbols),34 self.hps.data.filter_length // 2 + 1,35 self.hps.train.segment_size //36 self.hps.data.hop_length,37 midi_start=-5,38 midi_end=75,39 octave_range=24,40 n_speakers=len(self.hps.data.speakers),41 **self.hps.model42 ).to(self.device)43 44 _ = self.net_g.eval()45 _ = utils.load_checkpoint(args.checkpoint_path, model_g=self.net_g)46 self.interface = self._gradio_interface()47 48 def get_phoneme(self, text):49 cleaned_text, lang = clean_text(text)50 text_norm = cleaned_text_to_sequence(cleaned_text)51 52 if self.hps.data.add_blank:53 text_norm, lang = commons.intersperse_with_language_id(text_norm, lang, 0)54 55 text_norm = torch.LongTensor(text_norm)56 lang = torch.LongTensor(lang)57 58 return text_norm, lang, cleaned_text59 60 def inference(self, text, speaker_id_val, seed, scope_shift, duration):61 seed = int(seed)62 scope_shift = int(scope_shift)63 torch.manual_seed(seed)64 text_norm, tone, phones = self.get_phoneme(text)65 x_tst = text_norm.to(self.device).unsqueeze(0)66 t_tst = tone.to(self.device).unsqueeze(0)67 x_tst_lengths = torch.LongTensor([text_norm.size(0)]).to(self.device)68 speaker_id = torch.LongTensor([speaker_id_val]).to(self.device)69 70 decoder_inputs, *_ = self.net_g.infer_pre_decoder(71 x_tst,72 t_tst,73 x_tst_lengths,74 sid=speaker_id,75 noise_scale=0.667,76 noise_scale_w=0.8,77 length_scale=duration,78 scope_shift=scope_shift79 )80 81 audio = self.net_g.infer_decode_chunk(82 decoder_inputs, sid=speaker_id83 )[0, 0].data.cpu().float().numpy()84 85 del decoder_inputs,86 87 return phones, (self.hps.data.sampling_rate, audio)88 89 def _gradio_interface(self):90 title = "9Nine - PITS"91 92 self.inputs = [93 gr.Textbox(94 label="Text (150 words limitation)",95 value="[JA]そんなわけないじゃない。どうしてこうなるだろう。始めて好きな人ができた。一生ものの友达ができた。嬉しいことが二つ重なて。"96 "その二つの嬉しさがまたたくさんの嬉しさをつれて来てくれて。梦のように幸せの时间を手に入れたはずなのに。なのにどうして、こうなちょうだろう。[JA]",97 elem_id="tts-input"98 ),99 gr.Dropdown(100 list(self.hps.data.speakers),101 value=self.hps.data.speakers[1],102 label="Speaker Identity",103 type="index"104 ),105 gr.Slider(106 0, 65536, value=0, step=1, label="random seed"107 ),108 gr.Slider(109 -15, 15, value=0, step=1, label="scope-shift"110 ),111 gr.Slider(112 0.5, 2., value=1., step=0.1, label="duration multiplier"113 ),114 ]115 116 self.outputs = [117 gr.Textbox(label="Phonemes"),118 gr.Audio(type="numpy", label="Output audio")119 ]120 121 description = "9Nine - PITS"122 article = "Github: https://github.com/Aloento/VariTTS"123 examples = [["[JA]こんにちは、私は綾地寧々です。[JA]"]]124 125 return gr.Interface(126 fn=self.inference,127 inputs=self.inputs,128 outputs=self.outputs,129 title=title,130 description=description,131 article=article,132 cache_examples=False,133 examples=examples,134 )135 136 def launch(self):137 return self.interface.launch(share=False)138 139 140def parsearg():141 parser = argparse.ArgumentParser()142 parser.add_argument(143 '-c',144 '--config',145 type=str,146 default="./configs/config_cje.yaml",147 help='Path to configuration file'148 )149 parser.add_argument(150 '-m',151 '--model',152 type=str,153 default='9Nine',154 help='Model name'155 )156 parser.add_argument(157 '-r',158 '--checkpoint_path',159 type=str,160 default='./9Nine_Eval_71200.pth',161 help='Path to checkpoint for resume'162 )163 parser.add_argument(164 '-f',165 '--force_resume',166 type=str,167 help='Path to checkpoint for force resume'168 )169 parser.add_argument(170 '-d',171 '--dir',172 type=str,173 default='/DATA/audio/pits_samples',174 help='root dir'175 )176 args = parser.parse_args()177 return args178 179 180if __name__ == "__main__":181 import nltk182 nltk.download('cmudict')183 184 args = parsearg()185 app = GradioApp(args)186 app.launch()187 