CoolFace
Apppublic

aiqtech/SoulX-Singer

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
pipeline.py147 linesDownload Raw Back to preprocess
1import json2import shutil3import soundfile as sf4from pathlib import Path5import librosa6 7from preprocess.utils import convert_metadata, merge_short_segments8 9from preprocess.tools import (10    F0Extractor,11    VocalDetector,12    VocalSeparator,13    NoteTranscriber,14    LyricTranscriber,15)16 17 18class PreprocessPipeline:19    def __init__(self, device: str, language: str, save_dir: str, vocal_sep: bool = True, max_merge_duration: int = 60000):20        self.device = device21        self.language = language22        self.save_dir = save_dir23        self.vocal_sep = vocal_sep24        self.max_merge_duration = max_merge_duration25 26        if vocal_sep:27            self.vocal_separator = VocalSeparator(28                sep_model_path="pretrained_models/SoulX-Singer-Preprocess/mel-band-roformer-karaoke/mel_band_roformer_karaoke_becruily.ckpt",29                sep_config_path="pretrained_models/SoulX-Singer-Preprocess/mel-band-roformer-karaoke/config_karaoke_becruily.yaml",30                der_model_path="pretrained_models/SoulX-Singer-Preprocess/dereverb_mel_band_roformer/dereverb_mel_band_roformer_anvuew_sdr_19.1729.ckpt",31                der_config_path="pretrained_models/SoulX-Singer-Preprocess/dereverb_mel_band_roformer/dereverb_mel_band_roformer_anvuew.yaml",32                device=device33            )34        else:35            self.vocal_separator = None36        self.f0_extractor = F0Extractor(37            model_path="pretrained_models/SoulX-Singer-Preprocess/rmvpe/rmvpe.pt",38            device=device,39        )40        self.vocal_detector = VocalDetector(41            cut_wavs_output_dir=  f"{save_dir}/cut_wavs",42        )43        self.lyric_transcriber = LyricTranscriber(44            zh_model_path="pretrained_models/SoulX-Singer-Preprocess/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch",45            en_model_path="pretrained_models/SoulX-Singer-Preprocess/parakeet-tdt-0.6b-v2/parakeet-tdt-0.6b-v2.nemo",46            device=device47        )48        self.note_transcriber = NoteTranscriber(49            rosvot_model_path="pretrained_models/SoulX-Singer-Preprocess/rosvot/rosvot/model.pt", 50            rwbd_model_path="pretrained_models/SoulX-Singer-Preprocess/rosvot/rwbd/model.pt", 51            device=device52        )53 54    def run(55        self,56        audio_path: str,57        vocal_sep: bool = True,58        max_merge_duration: int = 60000,59        language: str = "Mandarin"60    ) -> None:61        vocal_sep = self.vocal_sep if vocal_sep is None else vocal_sep62        max_merge_duration = self.max_merge_duration if max_merge_duration is None else max_merge_duration63        language = self.language if language is None else language64        output_dir = Path(self.save_dir)65        output_dir.mkdir(parents=True, exist_ok=True)66 67        if vocal_sep:68            # Perform vocal/accompaniment separation69            sep = self.vocal_separator.process(audio_path)70            vocal = sep.vocals_dereverbed.T71            acc = sep.accompaniment.T72            sample_rate = sep.sample_rate73 74            vocal_path = output_dir / "vocal.wav"75            acc_path = output_dir / "acc.wav"76            sf.write(vocal_path, vocal, sample_rate)77            sf.write(acc_path, acc, sample_rate)78        else:79            # Use the original audio as vocal source (no separation)80            vocal, sample_rate = librosa.load(audio_path, sr=None, mono=True)81            vocal_path = output_dir / "vocal.wav"82            sf.write(vocal_path, vocal, sample_rate)83 84        vocal_f0 = self.f0_extractor.process(str(vocal_path))85        segments = self.vocal_detector.process(str(vocal_path), f0=vocal_f0)86 87        metadata = []88        for seg in segments:89            self.f0_extractor.process(seg["wav_fn"], f0_path=seg["wav_fn"].replace(".wav", "_f0.npy"))90            words, durs = self.lyric_transcriber.process(91                seg["wav_fn"], language92            )93            seg["words"] = words94            seg["word_durs"] = durs95            seg["language"] = language96            metadata.append(97                self.note_transcriber.process(seg, segment_info=seg)98            )99 100        merged = merge_short_segments(101            vocal,102            sample_rate,103            metadata,104            output_dir / "long_cut_wavs",105            max_duration_ms=max_merge_duration,106        )107 108        final_metadata = []109 110        for item in merged:111            self.f0_extractor.process(item.wav_fn, f0_path=item.wav_fn.replace(".wav", "_f0.npy"))112            final_metadata.append(convert_metadata(item))113 114        with open(output_dir / "metadata.json", "w", encoding="utf-8") as f:115            json.dump(final_metadata, f, ensure_ascii=False, indent=2)116 117        shutil.copy(output_dir / "metadata.json", audio_path.replace(".wav", ".json").replace(".mp3", ".json").replace(".flac", ".json"))118 119 120def main(args):121    pipeline = PreprocessPipeline(122        device=args.device,123        language=args.language,124        save_dir=args.save_dir,125        vocal_sep=args.vocal_sep,126        max_merge_duration=args.max_merge_duration,127    )128    pipeline.run(129        audio_path=args.audio_path,130        language=args.language131    )132 133 134if __name__ == "__main__":135    import argparse136 137    parser = argparse.ArgumentParser()138    parser.add_argument("--audio_path", type=str, required=True, help="Path to the input audio file")139    parser.add_argument("--save_dir", type=str, required=True, help="Directory to save the output files")140    parser.add_argument("--language", type=str, default="Mandarin", help="Language of the audio")141    parser.add_argument("--device", type=str, default="cuda:0", help="Device to run the models on")142    parser.add_argument("--vocal_sep", type=bool, default=True, help="Whether to perform vocal separation")143    parser.add_argument("--max_merge_duration", type=int, default=60000, help="Maximum merged segment duration in milliseconds")    144    args = parser.parse_args()145 146    main(args)147