Instructions to use SPRINGLab/SPRING_F5 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use SPRINGLab/SPRING_F5 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-to-speech", model="SPRINGLab/SPRING_F5", trust_remote_code=True)# Load model directly from transformers import SPRING_F5 model = SPRING_F5.from_pretrained("SPRINGLab/SPRING_F5", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
| import sys | |
| import os | |
| current_dir = os.path.dirname(os.path.abspath(__file__)) | |
| sys.path.append(current_dir) | |
| from transformers import PreTrainedModel, PretrainedConfig | |
| import torch | |
| import numpy as np | |
| from f5_tts.infer.utils_infer import ( | |
| infer_process, | |
| load_model, | |
| load_vocoder, | |
| preprocess_ref_audio_text, | |
| ) | |
| from f5_tts.model import DiT | |
| import soundfile as sf | |
| import io | |
| from pydub import AudioSegment, silence | |
| from huggingface_hub import hf_hub_download | |
| import os | |
| class SPRING_F5Config(PretrainedConfig): | |
| model_type = "SPRING_F5" | |
| def __init__(self, ckpt_path: str = "checkpoints/model_170000.pt", vocab_path: str = "checkpoints/vocab.txt", | |
| speed: float = 1.0, remove_sil: bool = True, **kwargs): | |
| super().__init__(**kwargs) | |
| self.ckpt_path = ckpt_path | |
| self.vocab_path = vocab_path | |
| self.speed = speed | |
| self.remove_sil = remove_sil | |
| class SPRING_F5Model(PreTrainedModel): | |
| config_class = SPRING_F5Config | |
| def __init__(self, config): | |
| super().__init__(config) | |
| self._device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| # Load vocoder | |
| self.vocoder = load_vocoder(vocoder_name="vocos", is_local=False, device=self._device) | |
| ckpt_file = hf_hub_download( repo_id=config.name_or_path, filename=config.ckpt_path) | |
| vocab_path = hf_hub_download(repo_id=config.name_or_path, filename=config.vocab_path ) | |
| self.ema_model = load_model( | |
| DiT, | |
| dict(dim=1024, depth=22, heads=16, ff_mult=2, text_dim=512, conv_layers=4), | |
| ckpt_file, | |
| mel_spec_type="vocos", | |
| vocab_file=vocab_path, | |
| device=self._device | |
| ) | |
| def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs): | |
| config = kwargs.pop("config", None) | |
| if config is None: | |
| config = SPRING_F5Config.from_pretrained(pretrained_model_name_or_path, **kwargs) | |
| config.name_or_path = pretrained_model_name_or_path | |
| return cls(config) | |
| def forward(self, text: str, ref_audio_path: str, ref_text: str, lang: str): | |
| """ | |
| Generate speech given a reference audio & text input. | |
| Args: | |
| text (str): The text to be synthesized. | |
| ref_audio_path (str): Path to the reference audio file. | |
| ref_text (str): The reference text. | |
| Returns: | |
| np.array: Generated waveform. | |
| """ | |
| if not os.path.exists(ref_audio_path): | |
| raise FileNotFoundError(f"Reference audio file {ref_audio_path} not found.") | |
| # Load reference audio & text | |
| ref_audio, ref_text = preprocess_ref_audio_text(ref_audio_path, ref_text) | |
| self.ema_model.to(self._device) | |
| self.vocoder.to(self._device) | |
| # Perform inference | |
| audio, final_sample_rate, _ = infer_process( | |
| ref_audio, | |
| ref_text, | |
| text, | |
| self.ema_model, | |
| self.vocoder, | |
| mel_spec_type="vocos", | |
| speed=self.config.speed, | |
| device=self._device, | |
| lang=lang # Language ID is used for number-to-Indic word conversion. | |
| ) | |
| # Convert to pydub format and remove silence if needed | |
| buffer = io.BytesIO() | |
| sf.write(buffer, audio, samplerate=24000, format="WAV") | |
| buffer.seek(0) | |
| audio_segment = AudioSegment.from_file(buffer, format="wav") | |
| if self.config.remove_sil: | |
| non_silent_segs = silence.split_on_silence( | |
| audio_segment, | |
| min_silence_len=1000, | |
| silence_thresh=-50, | |
| keep_silence=500, | |
| seek_step=10, | |
| ) | |
| non_silent_wave = sum(non_silent_segs, AudioSegment.silent(duration=0)) | |
| audio_segment = non_silent_wave | |
| # Normalize loudness | |
| target_dBFS = -20.0 | |
| change_in_dBFS = target_dBFS - audio_segment.dBFS | |
| audio_segment = audio_segment.apply_gain(change_in_dBFS) | |
| return np.array(audio_segment.get_array_of_samples()) | |