import json import math import os import unicodedata import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel from .configuration_yoruba_cfm import YorubaCFMConfig def sinusoidal_time_embedding(t, dim): half = dim // 2 device = t.device freqs = torch.exp( -math.log(10000) * torch.arange(0, half, device=device).float() / half ) args = t[:, None] * freqs[None, :] emb = torch.cat([torch.sin(args), torch.cos(args)], dim=-1) if dim % 2 == 1: emb = F.pad(emb, (0, 1)) return emb class TextEncoder(nn.Module): def __init__(self, config): super().__init__() self.pad_token_id = config.pad_token_id self.token_emb = nn.Embedding( config.vocab_size, config.d_model, padding_idx=config.pad_token_id ) self.pos_emb = nn.Embedding(config.text_max_positions, config.d_model) enc_layer = nn.TransformerEncoderLayer( d_model=config.d_model, nhead=config.n_heads, dim_feedforward=config.d_model * 4, dropout=config.dropout, batch_first=True, norm_first=True, activation="gelu", ) self.encoder = nn.TransformerEncoder(enc_layer, num_layers=config.text_layers) def forward(self, phoneme_ids): B, L = phoneme_ids.shape pos = torch.arange(L, device=phoneme_ids.device).unsqueeze(0).expand(B, L) x = self.token_emb(phoneme_ids) + self.pos_emb(pos) pad_mask = phoneme_ids.eq(self.pad_token_id) x = self.encoder(x, src_key_padding_mask=pad_mask) return x, pad_mask class DiTBlock(nn.Module): def __init__(self, config): super().__init__() self.self_attn = nn.MultiheadAttention( embed_dim=config.d_model, num_heads=config.n_heads, batch_first=True, dropout=config.dropout, ) self.cross_attn = nn.MultiheadAttention( embed_dim=config.d_model, num_heads=config.n_heads, batch_first=True, dropout=config.dropout, ) self.ffn = nn.Sequential( nn.Linear(config.d_model, config.d_model * 4), nn.GELU(), nn.Linear(config.d_model * 4, config.d_model), nn.Dropout(config.dropout), ) self.norm1 = nn.LayerNorm(config.d_model) self.norm2 = nn.LayerNorm(config.d_model) self.norm3 = nn.LayerNorm(config.d_model) def forward(self, x, cond, cond_mask=None): h = self.norm1(x) x = x + self.self_attn(h, h, h, need_weights=False)[0] h = self.norm2(x) x = x + self.cross_attn( h, cond, cond, key_padding_mask=cond_mask, need_weights=False )[0] x = x + self.ffn(self.norm3(x)) return x class LatentCFMDiT(nn.Module): def __init__(self, config): super().__init__() self.latent_in = nn.Linear(config.latent_dim, config.d_model) self.latent_out = nn.Linear(config.d_model, config.latent_dim) self.latent_pos = nn.Embedding(config.max_len, config.d_model) self.text_encoder = TextEncoder(config) self.time_mlp = nn.Sequential( nn.Linear(config.d_model, config.d_model * 4), nn.GELU(), nn.Linear(config.d_model * 4, config.d_model), ) self.blocks = nn.ModuleList( [DiTBlock(config) for _ in range(config.n_layers)] ) self.final_norm = nn.LayerNorm(config.d_model) def forward(self, x, t, phoneme_ids): B, T, _ = x.shape cond, cond_mask = self.text_encoder(phoneme_ids) pos = torch.arange(T, device=x.device).unsqueeze(0).expand(B, T) h = self.latent_in(x) + self.latent_pos(pos) time_emb = sinusoidal_time_embedding(t, h.shape[-1]) time_emb = self.time_mlp(time_emb).unsqueeze(1) h = h + time_emb for block in self.blocks: h = block(h, cond, cond_mask) h = self.final_norm(h) return self.latent_out(h) class YorubaCFMForTTS(PreTrainedModel): config_class = YorubaCFMConfig def __init__(self, config): super().__init__(config) self.cfm = LatentCFMDiT(config) self._encodec = None self._g2p = None self._phoneme_to_id = None self.post_init() def _load_phoneme_vocab(self): if self._phoneme_to_id is not None: return model_dir = self.config._name_or_path local_path = os.path.join(model_dir, "phoneme_vocab.json") if os.path.isfile(local_path): with open(local_path, encoding="utf-8") as f: vocab = json.load(f) self._phoneme_to_id = vocab["phoneme_to_id"] return from huggingface_hub import hf_hub_download path = hf_hub_download( repo_id=model_dir, filename="phoneme_vocab.json" ) with open(path, encoding="utf-8") as f: vocab = json.load(f) self._phoneme_to_id = vocab["phoneme_to_id"] def _load_encodec(self): if self._encodec is not None: return from transformers import EncodecModel self._encodec = EncodecModel.from_pretrained( self.config.encodec_model_id ) self._encodec.to(self.device).eval() for p in self._encodec.parameters(): p.requires_grad = False def _load_g2p(self): if self._g2p is not None: return try: from yoruba_g2p import YorubaG2P except ImportError: raise ImportError( "yoruba-g2p is required for text input. " "Install it with: pip install yoruba-g2p\n" "Alternatively, pass phoneme_ids directly." ) self._g2p = YorubaG2P() def _text_to_phoneme_ids(self, text): self._load_g2p() self._load_phoneme_vocab() text = unicodedata.normalize("NFC", text).strip() result = self._g2p.yoruba_word_to_ipa_phones(text) if isinstance(result, tuple): result = result[0] if isinstance(result, list): tokens = [t for t in result if t.strip()] else: tokens = result.strip().split() bos = self.config.bos_token_id eos = self.config.eos_token_id unk = self.config.unk_token_id ids = [bos] ids.extend(self._phoneme_to_id.get(tok, unk) for tok in tokens) ids.append(eos) return torch.LongTensor(ids).unsqueeze(0) @torch.no_grad() def generate( self, text=None, phoneme_ids=None, num_latent_frames=None, num_ode_steps=None, ): if text is None and phoneme_ids is None: raise ValueError("Provide either `text` or `phoneme_ids`.") if phoneme_ids is None: phoneme_ids = self._text_to_phoneme_ids(text).to(self.device) elif phoneme_ids.dim() == 1: phoneme_ids = phoneme_ids.unsqueeze(0) phoneme_ids = phoneme_ids.to(self.device) T = num_latent_frames or self.config.default_target_length D = self.config.latent_dim steps = num_ode_steps or self.config.num_ode_steps self.cfm.eval() x = torch.randn(1, T, D, device=self.device) dt = 1.0 / steps for step in range(steps): t = torch.full((1,), step / steps, device=self.device) v = self.cfm(x, t, phoneme_ids) x = x + dt * v self._load_encodec() latents = x.transpose(1, 2).contiguous() audio = self._encodec.decoder(latents) if isinstance(audio, (tuple, list)): audio = audio[0] return {"audio": audio, "sample_rate": self.config.sample_rate}