Download gooya_tts/runtime.py from Reza2kn/Gooya-RizehPizeh-v2-exp: direct link, hf CLI and curl.
- Browser
- Download file 6.81 kB
-
https://huggingface.co/Reza2kn/Gooya-RizehPizeh-v2-exp/resolve/main/gooya_tts/runtime.py
- Command line
-
hf download hf://Reza2kn/Gooya-RizehPizeh-v2-exp/gooya_tts/runtime.py
-
curl -L -o runtime.py https://huggingface.co/Reza2kn/Gooya-RizehPizeh-v2-exp/resolve/main/gooya_tts/runtime.py
6.81 kB
| """Self-contained inference for Gooya v2 experimental releases.""" | |
| from pathlib import Path | |
| import json,re,warnings | |
| import numpy as np | |
| import torch | |
| from torch import nn | |
| from torch.nn import functional as F | |
| from torch.nn.utils.rnn import pack_padded_sequence,pad_packed_sequence | |
| from torch.nn.utils import weight_norm | |
| from safetensors.torch import load_file | |
| from transformers import AutoTokenizer,AutoModelForSeq2SeqLM | |
| from .networks import PhonemeEncoder,MelDecoder,Phoneme2Mel | |
| from .hifigan import Generator | |
| from .normalize import normalize_persian_for_g2p | |
| from .words import surface_words | |
| class ContextEncoder(nn.Module): | |
| def __init__(self,local): | |
| super().__init__(); self.local=local; self.context=nn.GRU(80,128,batch_first=True,bidirectional=True); self.project=nn.Linear(256,32); self.dropout=nn.Dropout(.1) | |
| def forward(self,phones): | |
| x=self.local.embed(phones); lengths=(phones!=0).sum(-1).cpu() | |
| packed=pack_padded_sequence(x,lengths,batch_first=True,enforce_sorted=False) | |
| y,_=self.context(packed); y,_=pad_packed_sequence(y,batch_first=True,total_length=phones.shape[1]) | |
| return self.local(phones)+self.project(self.dropout(y)) | |
| class SmoothDuration(nn.Module): | |
| def __init__(self,original): | |
| super().__init__(); self.original=original; self.original.duration=False | |
| def forward(self,x): return F.softplus(self.original(x))+1 | |
| def get_embedding(self,*args,**kwargs): return self.original.get_embedding(*args,**kwargs) | |
| class AttrDict(dict): | |
| __getattr__=dict.__getitem__ | |
| class GooyaTTS: | |
| """Loads all model weights locally. CPU is the default; CUDA is optional.""" | |
| def __init__(self,model_dir,device='cpu',num_threads=4): | |
| self.root=Path(model_dir); self.device=torch.device(device) | |
| if num_threads: torch.set_num_threads(num_threads) | |
| self.config=json.loads((self.root/'gooya_config.json').read_text()); self.sample_rate=self.config['sample_rate']; self.ids=self.config['vocabulary'] | |
| self.tokenizer=AutoTokenizer.from_pretrained(self.root/'frontend',local_files_only=True) | |
| self.g2p=AutoModelForSeq2SeqLM.from_pretrained(self.root/'frontend',local_files_only=True).to(self.device).eval() | |
| stats=self.config['stats']; encoder=PhonemeEncoder(pitch_stats=stats['pitch'][:2],energy_stats=stats['energy'][:2]); decoder=MelDecoder() | |
| encoder.encoder.embed=nn.Embedding(len(self.ids),80,padding_idx=0) | |
| encoder.boundary_token_ids=[i for s,i in self.ids.items() if s.startswith('<')] | |
| encoder.encoder=ContextEncoder(encoder.encoder); encoder.duration_decoder=SmoothDuration(encoder.duration_decoder) | |
| decoder.mel_linear_down=nn.Linear(320,100); self.acoustic=Phoneme2Mel(encoder,decoder) | |
| self.acoustic.load_state_dict(load_file(str(self.root/'acoustic.safetensors')),strict=True); self.acoustic.to(self.device).eval() | |
| self.vocoder=Generator(AttrDict(self.config['vocoder_config'])); self.vocoder.conv_pre=weight_norm(nn.Conv1d(100,128,7,1,padding=3)) | |
| self.vocoder.load_state_dict(load_file(str(self.root/'vocoder.safetensors')),strict=True); self.vocoder.to(self.device).eval() | |
| self.overlay={} | |
| if self.config['frontend']['overlay']: | |
| data=json.loads((self.root/'frontend/overlay.json').read_text()); self.overlay={(r['surface'],r['raw']):r['target'] for r in data['rules']} | |
| count=sum(p.numel() for m in [self.g2p,self.acoustic,self.vocoder] for p in m.parameters()) | |
| if count!=self.config['parameters']['total']: raise ValueError(f'Parameter count mismatch: {count}') | |
| def from_pretrained(cls,repo_id,revision=None,**kwargs): | |
| from huggingface_hub import snapshot_download | |
| return cls(snapshot_download(repo_id,revision=revision),**kwargs) | |
| def synthesize(self,text,return_details=False): | |
| if not isinstance(text,str) or not text.strip(): raise ValueError('Text must be a nonempty string.') | |
| chunks=re.findall(r'[^ุ,ุ?\.\n]+[ุ,ุ?\.]?',text); waves=[]; records=[] | |
| for chunk in chunks: | |
| chunk=chunk.strip() | |
| if not chunk: continue | |
| normalized=normalize_persian_for_g2p(chunk) | |
| x=self.tokenizer(normalized,return_tensors='pt',add_special_tokens=False).to(self.device) | |
| if x['input_ids'].shape[1]>1024: raise ValueError('Clause exceeds 1024 input tokens; use shorter sentences or punctuation.') | |
| output=self.g2p.generate(**x,max_new_tokens=self.config['frontend']['max_new_tokens'],num_beams=self.config['frontend']['num_beams']) | |
| if self.tokenizer.eos_token_id not in output[0,1:].tolist(): raise RuntimeError('G2P did not terminate. Split this clause into shorter phrases; no incomplete audio was returned.') | |
| phones=self.tokenizer.decode(output[0],skip_special_tokens=True).strip().replace('?','Q') | |
| words=surface_words(normalized); pw=phones.split(); matched=len(words)==len(pw) | |
| if self.overlay and matched: phones=' '.join(self.overlay.get((a,b),b) for a,b in zip(words,pw)) | |
| if not matched: warnings.warn('G2P word count differs from input; inspect pronunciation in returned details.',RuntimeWarning) | |
| if not phones: raise RuntimeError('G2P returned an empty pronunciation.') | |
| tokens=['<sil>']+['<wb>' if c==' ' else c for c in phones]+['<comma>' if chunk[-1] in 'ุ,' else '<stop>'] | |
| unknown=set(tokens)-set(self.ids) | |
| if unknown: raise ValueError(f'Unsupported G2P symbols: {sorted(unknown)}') | |
| encoded=torch.tensor([[self.ids[t] for t in tokens]],device=self.device) | |
| pred=self.acoustic.encoder({'phoneme':encoded},train=False); mel=self.acoustic.decoder(pred['features']) | |
| if mel.shape[1]>24000//256*120: raise RuntimeError('Predicted clause duration exceeds 120 seconds.') | |
| wave=self.vocoder(F.pad(mel.transpose(1,2),(4,4),mode='replicate'))[:,:,1024:-1024].squeeze().cpu().numpy() | |
| if not np.isfinite(wave).all(): raise RuntimeError('Nonfinite audio output.') | |
| waves.append(wave); gap=self.config['comma_gap_seconds'] if chunk[-1] in 'ุ,' else self.config['sentence_gap_seconds']; waves.append(np.zeros(round(gap*self.sample_rate),dtype=np.float32)) | |
| records.append({'text':chunk,'normalized':normalized,'phonemes':phones,'word_count_matches':matched,'duration_seconds':len(wave)/self.sample_rate}) | |
| if not waves: raise ValueError('No speakable text found.') | |
| result=np.concatenate(waves[:-1]); details={'text':text,'model':self.config['name'],'parameters':self.config['parameters'],'sample_rate':self.sample_rate,'segments':records,'duration_seconds':len(result)/self.sample_rate} | |
| return (result,details) if return_details else result | |