Reza2kn's picture
Release self-contained experimental Persian TTS with verified offline CPU inference
f68f12f verified
Raw History Blame Contribute Delete
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}')
@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=['<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