CoolFace
Modelpublic

Aditya02/IndicF5

sourceHugging Facemitupdated 5mo agoView on Hugging Face
4likes428downloads
model.py246 linesDownload Raw Back to root
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())