Aditya02/IndicF5
4428
1import sys2import os3 4current_dir = os.path.dirname(os.path.abspath(__file__))5sys.path.append(current_dir)6 7from transformers import PreTrainedModel, PretrainedConfig, AutoConfig8import torch9import numpy as np10from f5_tts.infer.utils_infer import (11 infer_process,12 load_model,13 load_vocoder,14 preprocess_ref_audio_text,15)16from f5_tts.model import DiT17import soundfile as sf18import io19from pydub import AudioSegment, silence20from huggingface_hub import hf_hub_download21from safetensors.torch import load_file22import os23 24class INF5Config(PretrainedConfig):25 model_type = "inf5"26 27 def __init__(self, ckpt_path: str = "checkpoints/model_best.pt", vocab_path: str = "checkpoints/vocab.txt",28 speed: float = 1.0, remove_sil: bool = True, **kwargs):29 super().__init__(**kwargs)30 self.ckpt_path = ckpt_path31 self.vocab_path = vocab_path32 self.speed = speed33 self.remove_sil = remove_sil34 35class INF5Model(PreTrainedModel):36 config_class = INF5Config37 _tied_weights_keys = [] # Fix for transformers 5.0.0 compatibility38 39 @property40 def all_tied_weights_keys(self):41 """Compatibility property for transformers 5.0.0"""42 return {}43 44 45 def __init__(self, config):46 super().__init__(config)47 # Determine target device for inference (GPU if available)48 self._target_device = torch.device("cuda" if torch.cuda.is_available() else "cpu")49 50 # Disable torch.compile graph tracing to prevent ODE solver issues51 torch._dynamo.config.suppress_errors = True52 torch.backends.cudnn.deterministic = True53 torch.backends.cudnn.benchmark = False54 55 # Load vocoder - force on actual device to avoid meta tensor issues in transformers 5.0+56 with torch.device('cpu'):57 # Use eager backend to keep _orig_mod structure without actual compilation58 self.vocoder = torch.compile(load_vocoder(vocoder_name="vocos", is_local=False, device='cpu'), backend="eager")59 60 # Download and load model weights (load on CPU first for safe init,61 # model will be moved to target device in forward())62 safetensors_path = hf_hub_download(config.name_or_path, filename="model.safetensors")63 print(f"Loading model weights from {safetensors_path} (safetensors)...")64 state_dict = load_file(safetensors_path, device='cpu')65 66 # Download vocab.txt from HF Hub67 vocab_path = hf_hub_download(config.name_or_path, filename="checkpoints/vocab.txt")68 69 # Force model loading on CPU to avoid meta tensor issues70 with torch.device('cpu'):71 self.ema_model = load_model(72 DiT,73 dict(dim=1024, depth=22, heads=16, ff_mult=2, text_dim=512, conv_layers=4),74 mel_spec_type="vocos",75 vocab_file=vocab_path,76 device='cpu'77 )78 79 # Load state dict into model BEFORE compiling80 # Separate ema_model and vocoder weights, strip _orig_mod. prefix81 ema_state_dict = {}82 vocoder_state_dict = {}83 84 for key, value in state_dict.items():85 # Process ema_model weights86 if key.startswith("ema_model._orig_mod."):87 new_key = key.replace("ema_model._orig_mod.", "")88 ema_state_dict[new_key] = value89 elif key.startswith("ema_model."):90 new_key = key.replace("ema_model.", "")91 ema_state_dict[new_key] = value92 # Process vocoder weights93 elif key.startswith("vocoder._orig_mod."):94 new_key = key.replace("vocoder._orig_mod.", "")95 vocoder_state_dict[new_key] = value96 elif key.startswith("vocoder."):97 new_key = key.replace("vocoder.", "")98 vocoder_state_dict[new_key] = value99 100 # Load ema_model weights101 missing_keys, unexpected_keys = self.ema_model.load_state_dict(102 ema_state_dict, strict=False)103 104 # Load vocoder weights if any (vocoder is already compiled, so use _orig_mod if needed)105 if vocoder_state_dict:106 try:107 # Try loading directly first108 self.vocoder.load_state_dict(vocoder_state_dict, strict=False)109 except:110 # If vocoder is compiled, access the underlying model111 if hasattr(self.vocoder, '_orig_mod'):112 self.vocoder._orig_mod.load_state_dict(vocoder_state_dict, strict=False)113 114 # Use eager backend - disables actual compilation while keeping _orig_mod115 # structure for weight serialization. Full torch.compile with inductor116 # breaks the ODE solver in CFM.sample() causing jumbled/partial text output.117 self.ema_model = torch.compile(self.ema_model, backend="eager")118 print(f"Weight loading - Missing keys: {len(missing_keys)}, Unexpected keys: {len(unexpected_keys)}")119 if missing_keys:120 print(f"Missing keys sample: {missing_keys[:5]}")121 if unexpected_keys:122 print(f"Unexpected keys sample: {unexpected_keys[:5]}")123 124 # Flag for lazy buffer recomputation (see _recompute_buffers).125 # We cannot recompute here because transformers 5.0 materializes126 # meta tensors AFTER __init__ returns, overwriting our values.127 self._buffers_need_recompute = True128 129 def _recompute_buffers(self):130 """Recompute non-persistent buffers that were corrupted by131 transformers 5.0's meta device initialization.132 133 transformers 5.0 wraps __init__ in torch.device('meta') context,134 then materializes meta tensors with uninitialized (garbage) values.135 Non-persistent buffers (not in safetensors) never get correct values.136 This method must be called AFTER from_pretrained completes."""137 from f5_tts.model.modules import precompute_freqs_cis138 139 # Get the underlying model (unwrap torch.compile if needed)140 ema = self.ema_model._orig_mod if hasattr(self.ema_model, '_orig_mod') else self.ema_model141 142 # Determine current device of the buffers143 buf_device = ema.transformer.text_embed.freqs_cis.device if (144 hasattr(ema, 'transformer') and hasattr(ema.transformer, 'text_embed')145 and hasattr(ema.transformer.text_embed, 'freqs_cis')146 ) else torch.device('cpu')147 148 # Recompute text_embed.freqs_cis (positional embeddings for text)149 if hasattr(ema, 'transformer') and hasattr(ema.transformer, 'text_embed'):150 text_embed = ema.transformer.text_embed151 if hasattr(text_embed, 'extra_modeling') and text_embed.extra_modeling:152 text_dim = text_embed.text_embed.embedding_dim153 max_pos = text_embed.precompute_max_pos154 freqs_cis = precompute_freqs_cis(text_dim, max_pos).to(buf_device)155 # Check if recomputation needed (first value should be cos(0) = 1.0)156 if text_embed.freqs_cis.is_meta or abs(text_embed.freqs_cis[0, 0].item() - 1.0) > 0.01:157 text_embed.freqs_cis.data.copy_(freqs_cis)158 print(f"Recomputed freqs_cis: shape={freqs_cis.shape}, first_val={freqs_cis[0,0].item():.4f}")159 160 # Recompute mel_spec.dummy buffer161 if hasattr(ema, 'mel_spec') and hasattr(ema.mel_spec, 'dummy'):162 if ema.mel_spec.dummy.is_meta or ema.mel_spec.dummy.item() != 0:163 ema.mel_spec.dummy.data.fill_(0)164 print("Recomputed mel_spec.dummy to 0")165 166 # Recompute rotary_embed.inv_freq if needed167 if hasattr(ema, 'transformer') and hasattr(ema.transformer, 'rotary_embed'):168 rot = ema.transformer.rotary_embed169 if hasattr(rot, 'inv_freq'):170 dim = rot.inv_freq.shape[0] * 2171 if rot.inv_freq.is_meta or rot.inv_freq[0].abs() > 10:172 theta = 10000.0173 inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2).float().to(buf_device) / dim))174 rot.inv_freq.data.copy_(inv_freq)175 print(f"Recomputed rotary inv_freq: shape={inv_freq.shape}")176 177 self._buffers_need_recompute = False178 179 180 @property181 def device(self):182 """Get the target device of the model (GPU if available, else CPU)"""183 return getattr(self, '_target_device', torch.device('cpu'))184 185 def forward(self, text: str, ref_audio_path: str, ref_text: str):186 """187 Generate speech given a reference audio & text input.188 189 Args:190 text (str): The text to be synthesized.191 ref_audio_path (str): Path to the reference audio file.192 ref_text (str): The reference text.193 Returns:194 np.array: Generated waveform.195 """196 197 # Lazy recomputation of non-persistent buffers corrupted by transformers 5.0198 if getattr(self, "_buffers_need_recompute", False):199 self._recompute_buffers()200 201 if not os.path.exists(ref_audio_path):202 raise FileNotFoundError(f"Reference audio file {ref_audio_path} not found.")203 204 # Load reference audio & text205 ref_audio, ref_text = preprocess_ref_audio_text(ref_audio_path, ref_text)206 207 # Move models to target device (GPU if available) - only actually208 # transfers on first call; subsequent calls are no-ops209 self.ema_model.to(self.device)210 self.vocoder.to(self.device)211 212 # Perform inference213 audio, final_sample_rate, _ = infer_process(214 ref_audio,215 ref_text,216 text,217 self.ema_model,218 self.vocoder,219 mel_spec_type="vocos",220 speed=self.config.speed,221 device=self.device,222 )223 224 # Convert to pydub format and remove silence if needed225 buffer = io.BytesIO()226 sf.write(buffer, audio, samplerate=24000, format="WAV")227 buffer.seek(0)228 audio_segment = AudioSegment.from_file(buffer, format="wav")229 230 if self.config.remove_sil:231 non_silent_segs = silence.split_on_silence(232 audio_segment,233 min_silence_len=1000,234 silence_thresh=-50,235 keep_silence=500,236 seek_step=10,237 )238 non_silent_wave = sum(non_silent_segs, AudioSegment.silent(duration=0))239 audio_segment = non_silent_wave240 241 # Normalize loudness242 target_dBFS = -20.0243 change_in_dBFS = target_dBFS - audio_segment.dBFS244 audio_segment = audio_segment.apply_gain(change_in_dBFS)245 246 return np.array(audio_segment.get_array_of_samples())