CoolFace
Apppublic

Aloento/9Nine-PITS

sourceHugging Faceagpl-3.0updated 4y agoView on Hugging Face
1likes
app.py187 linesDownload Raw Back to root
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