"""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}') @classmethod def from_pretrained(cls,repo_id,revision=None,**kwargs): from huggingface_hub import snapshot_download return cls(snapshot_download(repo_id,revision=revision),**kwargs) @torch.inference_mode() 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=['']+['' if c==' ' else c for c in phones]+['' if chunk[-1] in '،,' else ''] 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