diff --git a/README.md b/README.md index 31f21b0..7f333fa 100644 --- a/README.md +++ b/README.md @@ -21,13 +21,15 @@ ## 🎵 Overview -**SoulX-Singer** is a high-fidelity, zero-shot singing voice synthesis model that enables users to generate realistic singing voices for unseen singers. -It supports **melody-conditioned (F0 contour)** and **score-conditioned (MIDI notes)** control for precise pitch, rhythm, and expression. +**SoulX-Singer** is a high-fidelity, zero-shot singing voice synthesis model that enables users to generate realistic singing voices for unseen singers. It supports **melody-conditioned (F0 contour)** and **score-conditioned (MIDI notes)** control for precise pitch, rhythm, and expression. + +**SoulX-Singer-SVC** is a singing voice conversion (SVC) model finetuned from **SoulX-Singer**. Singing Voice Conversion aims to transform a source singing recording into the target singer’s voice while preserving the original melody, rhythm, and lyrical content. Based on the strong generative capability of SoulX-Singer, SoulX-Singer-SVC enables high-quality singing voice conversion directly from raw singing audio, without requiring lyric or MIDI transcriptions. --- ## ✨ Key Features +#### SoulX-Singer - **🎤 Zero-Shot Singing** – Generate high-fidelity voices for unseen singers, no fine-tuning needed. - **🎵 Flexible Control Modes** – Melody (F0) and Score (MIDI) conditioning. - **📚 Large-Scale Dataset** – 42,000+ hours of aligned vocals, lyrics, notes across Mandarin, English, Cantonese. @@ -35,6 +37,11 @@ It supports **melody-conditioned (F0 contour)** and **score-conditioned (MIDI no - **✏️ Singing Voice Editing** – Modify lyrics while keeping natural prosody. - **🌐 Cross-Lingual Synthesis** – High-fidelity synthesis by disentangling timbre from content. +#### SoulX-Singer-SVC +- **🎙️ Zero-Shot Timbre and Style Transfer** – Transfer singer identity and style to unseen voices without per-speaker fine-tuning. +- **🌍 Language-Agnostic Conversion** – Works across multilingual singing content. +- **🔄 Transcription-Free Audio-to-Audio Conversion** – Convert target singing directly without lyrics transcription or MIDI inputs. + ---

@@ -45,7 +52,7 @@ It supports **melody-conditioned (F0 contour)** and **score-conditioned (MIDI no ## 🎬 Demo Examples - +### Singing Voice Synthesis (SVS)

@@ -57,9 +64,17 @@ It supports **melody-conditioned (F0 contour)** and **score-conditioned (MIDI no
+### Singing Voice Conversion (SVC) +
+ + + +
+ --- ## 📰 News +- **[2026-03-16]** [SoulX-Singer-SVC](https://huggingface.co/Soul-AILab/SoulX-Singer/blob/main/model-svc.pt) is released, and [SoulX-Singer Online Demo](https://huggingface.co/spaces/Soul-AILab/SoulX-Singer) has been updated to support singing voice conversion (SVC). - **[2026-02-12]** [SoulX-Singer Eval Dataset](https://huggingface.co/datasets/Soul-AILab/SoulX-Singer-Eval-Dataset) is now available on Hugging Face Datasets. - **[2026-02-09]** [SoulX-Singer Online Demo](https://huggingface.co/spaces/Soul-AILab/SoulX-Singer) is live on Hugging Face Spaces — try singing voice synthesis in your browser. - **[2026-02-08]** [MIDI Editor](https://huggingface.co/spaces/Soul-AILab/SoulX-Singer-Midi-Editor) is available on Hugging Face Spaces. @@ -105,11 +120,11 @@ Install Hugging Face Hub if needed: pip install -U huggingface_hub ``` -Download the SVS model and preprocessing models: +Download the SVS, SVC model and preprocessing models: ```sh pip install -U huggingface_hub -# Download the SoulX-Singer SVS model +# Download the SoulX-Singer SVS and SVC model hf download Soul-AILab/SoulX-Singer --local-dir pretrained_models/SoulX-Singer # Download models required for preprocessing @@ -119,7 +134,7 @@ hf download Soul-AILab/SoulX-Singer-Preprocess --local-dir pretrained_models/Sou ### 4. Run the Demo -Run the inference demo: +#### Run the SVS inference demo ``` sh bash example/infer.sh ``` @@ -132,14 +147,30 @@ The metadata produced by the automatic preprocessing pipeline may not perfectly How to use the Midi-Editor: - [Eiditing Metadata with Midi-Editor](preprocess/README.md#L104-L105) +#### Run the SVC inference demo + +```sh +bash example/infer_svc.sh +``` + +This example performs audio-to-audio SVC, converting the target singing into the prompt timbre using waveform and F0 inputs. +To prepare your own SVC data, run `example/preprocess.sh` with `midi_transcribe=False`. + + ### 🌐 WebUI -You can launch the interactive interface with: +You can launch the interactive interface for SVS (Synthesised from lyrics and MIDI transcriptions) with: ``` python webui.py ``` +For SVC WebUI (audio-to-audio conversion): + +``` +python webui_svc.py +``` + ## 🚧 Roadmap @@ -150,7 +181,7 @@ python webui.py - [x] 📊 Release the SoulX-Singer-Eval benchmark - [ ] 🎹 Inference support for user-friendly MIDI-based input - [ ] 📚 Comprehensive tutorials and usage documentation -- [ ] 🎵 Support for wav-to-wav singing voice conversion (without transcription) +- [x] 🎵 Support for wav-to-wav singing voice conversion (without transcription) ## 🙏 Acknowledgements diff --git a/assets/technical-report.pdf b/assets/technical-report.pdf index 05eadb6..cd85729 100644 Binary files a/assets/technical-report.pdf and b/assets/technical-report.pdf differ diff --git a/cli/inference.py b/cli/inference.py index d91926b..148cf73 100644 --- a/cli/inference.py +++ b/cli/inference.py @@ -17,6 +17,7 @@ def build_model( model_path: str, config: DictConfig, device: str = "cuda", + use_fp16: bool = False, ): """ Build the model from the pre-trained model path and model configuration. @@ -25,9 +26,10 @@ def build_model( model_path (str): Path to the checkpoint file. config (DictConfig): Model configuration. device (str, optional): Device to use. Defaults to "cuda". + use_fp16 (bool, optional): If True and device is CUDA, convert model to FP16 after load. Defaults to False. Returns: - Tuple[torch.nn.Module, torch.nn.Module]: The initialized model and vocoder. + SoulXSinger: The initialized model. """ if not os.path.isfile(model_path): @@ -39,7 +41,7 @@ def build_model( print("Model initialized.") print("Model parameters:", sum(p.numel() for p in model.parameters()) / 1e6, "M") - checkpoint = torch.load(model_path, weights_only=False, map_location=device) + checkpoint = torch.load(model_path, weights_only=False, map_location="cpu") if "state_dict" not in checkpoint: raise KeyError( f"Checkpoint at {model_path} has no 'state_dict' key. " @@ -47,6 +49,10 @@ def build_model( ) model.load_state_dict(checkpoint["state_dict"], strict=True) + if use_fp16 and ((isinstance(device, str) and device.startswith("cuda")) or (hasattr(device, "type") and getattr(device, "type", None) == "cuda")): + model.half() + model.mel.float() + print("Model converted to FP16 (mel kept in FP32).") model.eval() model.to(device) print("Model checkpoint loaded.") @@ -104,6 +110,7 @@ def process(args, config, model: torch.nn.Module): n_steps=config.infer.n_steps, cfg=config.infer.cfg, control=args.control, + use_fp16=args.use_fp16, ) generated_audio = generated_audio.squeeze().cpu().numpy() @@ -119,6 +126,7 @@ def main(args, config): model_path=args.model_path, config=config, device=args.device, + use_fp16=getattr(args, "use_fp16", False), ) process(args, config, model) @@ -141,7 +149,14 @@ if __name__ == "__main__": choices=["melody", "score"], help="Control mode: melody or score only", ) + parser.add_argument( + "--fp16", + action="store_true", + default=False, + help="Use FP16 inference (faster on GPU)", + ) args = parser.parse_args() - + args.use_fp16 = args.fp16 + config = load_config(args.config) main(args, config) diff --git a/cli/inference_svc.py b/cli/inference_svc.py new file mode 100644 index 0000000..8372b50 --- /dev/null +++ b/cli/inference_svc.py @@ -0,0 +1,130 @@ +import os +import torch +import json +import argparse +from tqdm import tqdm +import numpy as np +import soundfile as sf +from collections import OrderedDict +from omegaconf import DictConfig + +from soulxsinger.utils.file_utils import load_config +from soulxsinger.models.soulxsinger_svc import SoulXSingerSVC +from soulxsinger.utils.audio_utils import load_wav + + +def build_model( + model_path: str, + config: DictConfig, + device: str = "cuda", + use_fp16: bool = False, +): + """ + Build the model from the pre-trained model path and model configuration. + + Args: + model_path (str): Path to the checkpoint file. + config (DictConfig): Model configuration. + device (str, optional): Device to use. Defaults to "cuda". + use_fp16 (bool, optional): If True and device is CUDA, convert model to FP16 after load. Defaults to False. + + Returns: + SoulXSingerSVC: The initialized model. + """ + + if not os.path.isfile(model_path): + raise FileNotFoundError( + f"Model checkpoint not found: {model_path}. " + "Please download the pretrained model and place it at the path, or set --model_path." + ) + model = SoulXSingerSVC(config).to(device) + print("Model initialized.") + print("Model parameters:", sum(p.numel() for p in model.parameters()) / 1e6, "M") + + checkpoint = torch.load(model_path, weights_only=False, map_location="cpu") + if "state_dict" not in checkpoint: + raise KeyError( + f"Checkpoint at {model_path} has no 'state_dict' key. " + "Expected a checkpoint saved with model.state_dict()." + ) + model.load_state_dict(checkpoint["state_dict"], strict=True) + + if use_fp16 and ((isinstance(device, str) and device.startswith("cuda")) or (hasattr(device, "type") and getattr(device, "type", None) == "cuda")): + model.half() + model.mel.float() + print("Model converted to FP16 (mel kept in FP32).") + print("Model checkpoint loaded.") + model.eval() + model.to(device) + + return model + + +def process(args, config, model: torch.nn.Module): + """Run the full inference pipeline given a data_processor and model. + """ + + os.makedirs(args.save_dir, exist_ok=True) + pt_wav = load_wav(args.prompt_wav_path, config.audio.sample_rate).to(args.device) + gt_wav = load_wav(args.target_wav_path, config.audio.sample_rate).to(args.device) + pt_f0 = torch.from_numpy(np.load(args.prompt_f0_path)).unsqueeze(0).to(args.device) + gt_f0 = torch.from_numpy(np.load(args.target_f0_path)).unsqueeze(0).to(args.device) + + n_step = args.n_steps if hasattr(args, "n_steps") else config.infer.n_steps + cfg = args.cfg if hasattr(args, "cfg") else config.infer.cfg + + with torch.no_grad(): + generated_audio, generated_shift = model.infer( + pt_wav=pt_wav, + gt_wav=gt_wav, + pt_f0=pt_f0, + gt_f0=gt_f0, + auto_shift=args.auto_shift, + pitch_shift=args.pitch_shift, + n_steps=n_step, + cfg=cfg, + use_fp16=args.use_fp16, + ) + generated_audio = generated_audio.squeeze().float().cpu().numpy() + if args.pitch_shift != generated_shift: + args.pitch_shift = generated_shift + # print(f"Applied pitch shift of {generated_shift} semitones to match GT F0 contour.") + + sf.write(os.path.join(args.save_dir, "generated.wav"), generated_audio, config.audio.sample_rate) + print(f"Generated audio saved to {os.path.join(args.save_dir, 'generated.wav')}") + + +def main(args, config): + model = build_model( + model_path=args.model_path, + config=config, + device=args.device, + use_fp16=getattr(args, "use_fp16", False), + ) + process(args, config, model) + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--device", type=str, default="cuda") + parser.add_argument("--model_path", type=str, default='pretrained_models/soulx-singer/model.pt') + parser.add_argument("--config", type=str, default='soulxsinger/config/soulxsinger.yaml') + parser.add_argument("--prompt_wav_path", type=str, default='example/audio/zh_prompt.wav') + parser.add_argument("--target_wav_path", type=str, default='example/audio/zh_target.wav') + parser.add_argument("--prompt_f0_path", type=str, default='example/audio/zh_prompt_f0.npy') + parser.add_argument("--target_f0_path", type=str, default='example/audio/zh_target_f0.npy') + parser.add_argument("--save_dir", type=str, default='outputs') + parser.add_argument("--auto_shift", action="store_true") + parser.add_argument("--pitch_shift", type=int, default=0) + parser.add_argument("--n_steps", type=int, default=32) + parser.add_argument("--cfg", type=float, default=3.0) + parser.add_argument( + "--fp16", + action="store_true", + default=False, + help="Use FP16 inference (faster on GPU)", + ) + args = parser.parse_args() + args.use_fp16 = args.fp16 + + config = load_config(args.config) + main(args, config) diff --git a/example/audio/music_f0.npy b/example/audio/music_f0.npy new file mode 100644 index 0000000..3e48656 Binary files /dev/null and b/example/audio/music_f0.npy differ diff --git a/example/audio/svc_prompt_demo.mp3 b/example/audio/svc_prompt_demo.mp3 new file mode 100644 index 0000000..3c4757b Binary files /dev/null and b/example/audio/svc_prompt_demo.mp3 differ diff --git a/example/audio/svc_target_demo.mp3 b/example/audio/svc_target_demo.mp3 new file mode 100644 index 0000000..dc10124 Binary files /dev/null and b/example/audio/svc_target_demo.mp3 differ diff --git a/example/audio/zh_prompt_f0.npy b/example/audio/zh_prompt_f0.npy new file mode 100644 index 0000000..0308a4b Binary files /dev/null and b/example/audio/zh_prompt_f0.npy differ diff --git a/example/infer.sh b/example/infer.sh index a39661e..5f033c7 100644 --- a/example/infer.sh +++ b/example/infer.sh @@ -25,4 +25,5 @@ python -m cli.inference \ --phoneset_path $phoneset_path \ --save_dir $save_dir \ --auto_shift \ - --pitch_shift 0 \ No newline at end of file + --pitch_shift 0 \ + --fp16 \ No newline at end of file diff --git a/example/infer_svc.sh b/example/infer_svc.sh new file mode 100644 index 0000000..036ef83 --- /dev/null +++ b/example/infer_svc.sh @@ -0,0 +1,28 @@ +#!/bin/bash + +script_dir=$(dirname "$(realpath "$0")") +root_dir=$(dirname "$script_dir") + +cd $root_dir || exit +export PYTHONPATH=$root_dir:$PYTHONPATH + +model_path=pretrained_models/SoulX-Singer/model-svc.pt +config=soulxsinger/config/soulxsinger.yaml +prompt_wav_path=example/audio/zh_prompt.mp3 +target_wav_path=example/audio/music.mp3 +prompt_f0_path=example/audio/zh_prompt_f0.npy +target_f0_path=example/audio/music_f0.npy +save_dir=example/generated/music_svc + +python -m cli.inference_svc \ + --device cuda \ + --model_path $model_path \ + --config $config \ + --prompt_wav_path $prompt_wav_path \ + --target_wav_path $target_wav_path \ + --prompt_f0_path $prompt_f0_path \ + --target_f0_path $target_f0_path \ + --save_dir $save_dir \ + --auto_shift \ + --pitch_shift 0 \ + --fp16 \ No newline at end of file diff --git a/example/preprocess.sh b/example/preprocess.sh index bd6b4ce..0c8abf2 100644 --- a/example/preprocess.sh +++ b/example/preprocess.sh @@ -15,6 +15,7 @@ save_dir=example/transcriptions/zh_prompt language=Mandarin vocal_sep=False max_merge_duration=30000 +midi_transcribe=True # Whether to transcribe vocal midi, set True for singing voice synthesis, False for singing voice conversion python -m preprocess.pipeline \ --audio_path $audio_path \ @@ -22,7 +23,8 @@ python -m preprocess.pipeline \ --language $language \ --device $device \ --vocal_sep $vocal_sep \ - --max_merge_duration $max_merge_duration + --max_merge_duration $max_merge_duration \ + --midi_transcribe $midi_transcribe ####### Run Target Annotation ####### @@ -31,6 +33,7 @@ save_dir=example/transcriptions/music language=Mandarin vocal_sep=True max_merge_duration=60000 +midi_transcribe=True # Whether to transcribe vocal midi, set True for singing voice synthesis, False for singing voice conversion python -m preprocess.pipeline \ --audio_path $audio_path \ @@ -38,4 +41,5 @@ python -m preprocess.pipeline \ --language $language \ --device $device \ --vocal_sep $vocal_sep \ - --max_merge_duration $max_merge_duration \ No newline at end of file + --max_merge_duration $max_merge_duration \ + --midi_transcribe $midi_transcribe \ No newline at end of file diff --git a/preprocess/pipeline.py b/preprocess/pipeline.py index f3a69c7..7e5274b 100644 --- a/preprocess/pipeline.py +++ b/preprocess/pipeline.py @@ -16,12 +16,13 @@ from preprocess.tools import ( class PreprocessPipeline: - def __init__(self, device: str, language: str, save_dir: str, vocal_sep: bool = True, max_merge_duration: int = 60000): + def __init__(self, device: str, language: str, save_dir: str, vocal_sep: bool = True, max_merge_duration: int = 60000, midi_transcribe: bool = True): self.device = device self.language = language self.save_dir = save_dir self.vocal_sep = vocal_sep self.max_merge_duration = max_merge_duration + self.midi_transcribe = midi_transcribe if vocal_sep: self.vocal_separator = VocalSeparator( @@ -37,26 +38,31 @@ class PreprocessPipeline: model_path="pretrained_models/SoulX-Singer-Preprocess/rmvpe/rmvpe.pt", device=device, ) - self.vocal_detector = VocalDetector( - cut_wavs_output_dir= f"{save_dir}/cut_wavs", - ) - self.lyric_transcriber = LyricTranscriber( - zh_model_path="pretrained_models/SoulX-Singer-Preprocess/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch", - en_model_path="pretrained_models/SoulX-Singer-Preprocess/parakeet-tdt-0.6b-v2/parakeet-tdt-0.6b-v2.nemo", - device=device - ) - self.note_transcriber = NoteTranscriber( - rosvot_model_path="pretrained_models/SoulX-Singer-Preprocess/rosvot/rosvot/model.pt", - rwbd_model_path="pretrained_models/SoulX-Singer-Preprocess/rosvot/rwbd/model.pt", - device=device - ) + if self.midi_transcribe: + self.vocal_detector = VocalDetector( + cut_wavs_output_dir= f"{save_dir}/cut_wavs", + ) + self.lyric_transcriber = LyricTranscriber( + zh_model_path="pretrained_models/SoulX-Singer-Preprocess/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch", + en_model_path="pretrained_models/SoulX-Singer-Preprocess/parakeet-tdt-0.6b-v2/parakeet-tdt-0.6b-v2.nemo", + device=device + ) + self.note_transcriber = NoteTranscriber( + rosvot_model_path="pretrained_models/SoulX-Singer-Preprocess/rosvot/rosvot/model.pt", + rwbd_model_path="pretrained_models/SoulX-Singer-Preprocess/rosvot/rwbd/model.pt", + device=device + ) + else: + self.vocal_detector = None + self.lyric_transcriber = None + self.note_transcriber = None def run( self, audio_path: str, - vocal_sep: bool = True, - max_merge_duration: int = 60000, - language: str = "Mandarin" + vocal_sep: bool = None, + max_merge_duration: int = None, + language: str = None, ) -> None: vocal_sep = self.vocal_sep if vocal_sep is None else vocal_sep max_merge_duration = self.max_merge_duration if max_merge_duration is None else max_merge_duration @@ -81,7 +87,11 @@ class PreprocessPipeline: vocal_path = output_dir / "vocal.wav" sf.write(vocal_path, vocal, sample_rate) - vocal_f0 = self.f0_extractor.process(str(vocal_path)) + vocal_f0 = self.f0_extractor.process(str(vocal_path), f0_path=str(vocal_path).replace(".wav", "_f0.npy")) + + if not self.midi_transcribe or self.vocal_detector is None or self.lyric_transcriber is None or self.note_transcriber is None: + return + segments = self.vocal_detector.process(str(vocal_path), f0=vocal_f0) metadata = [] @@ -124,10 +134,11 @@ def main(args): save_dir=args.save_dir, vocal_sep=args.vocal_sep, max_merge_duration=args.max_merge_duration, + midi_transcribe=args.midi_transcribe, ) pipeline.run( audio_path=args.audio_path, - language=args.language + language=args.language, ) @@ -139,8 +150,12 @@ if __name__ == "__main__": parser.add_argument("--save_dir", type=str, required=True, help="Directory to save the output files") parser.add_argument("--language", type=str, default="Mandarin", help="Language of the audio") parser.add_argument("--device", type=str, default="cuda:0", help="Device to run the models on") - parser.add_argument("--vocal_sep", type=bool, default=True, help="Whether to perform vocal separation") + parser.add_argument("--vocal_sep", type=str, default="True", help="Whether to perform vocal separation") parser.add_argument("--max_merge_duration", type=int, default=60000, help="Maximum merged segment duration in milliseconds") + parser.add_argument("--midi_transcribe", type=str, default="True", help="Whether to do MIDI transcription") args = parser.parse_args() + args.vocal_sep = args.vocal_sep.lower() == "true" + args.midi_transcribe = args.midi_transcribe.lower() == "true" + main(args) diff --git a/preprocess/tools/vocal_separation/model.py b/preprocess/tools/vocal_separation/model.py index b4ee8bf..9da442f 100644 --- a/preprocess/tools/vocal_separation/model.py +++ b/preprocess/tools/vocal_separation/model.py @@ -54,7 +54,7 @@ def build_model(args): return model, config -def build_models(dict_args): +def build_models(dict_args, use_der: bool = False): args = parse_args_inference(dict_args) ########## load model ########## @@ -65,25 +65,26 @@ def build_models(dict_args): sep_model, sep_config = build_model(args) - args.config_path = args.der_config_path - args.start_check_point = args.der_start_check_point - - dereverb_model, dereverb_config = build_model(args) - - sep_model = sep_model - dereverb_model = dereverb_model + if use_der: + args.config_path = args.der_config_path + args.start_check_point = args.der_start_check_point + dereverb_model, dereverb_config = build_model(args) + else: + dereverb_model, dereverb_config = None, None return sep_model, sep_config, dereverb_model, dereverb_config, args def main(args, sep_model=None, sep_config=None, dereverb_model=None, dereverb_config=None, device=None): - ######## process data ########## sample_rate = getattr(sep_config.audio, 'sample_rate', 44100) path = args.input_path mix, _ = librosa.load(path, sr=sample_rate, mono=False) vocals = process(mix, sep_model, args, sep_config, device) - dereverbed_vocals = process(vocals.mean(0), dereverb_model, args, dereverb_config, device) + if dereverb_model is not None and dereverb_config is not None: + dereverbed_vocals = process(vocals.mean(0), dereverb_model, args, dereverb_config, device) + else: + dereverbed_vocals = vocals accompaniment = mix - dereverbed_vocals return mix, vocals, dereverbed_vocals, accompaniment, sample_rate @@ -113,6 +114,8 @@ class VocalSeparator: der_model_path: str, der_config_path: str, *, + chunk_length_sec: int = 5, + use_der: bool = False, model_type: str = "mel_band_roformer", disable_detailed_pbar: bool = True, device: str = "cuda", @@ -122,11 +125,14 @@ class VocalSeparator: Args: device: Torch device string, e.g. ``"cuda:0"``. + use_der: If True, load and run dereverb model; if False, skip dereverb (default False). model_type: Separation model type key. sep_config_path: Config path for separation model. sep_start_check_point: Checkpoint path for separation model. der_config_path: Config path for dereverb model. der_start_check_point: Checkpoint path for dereverb model. + chunk_length_sec: Chunk length in seconds. Set lower if you want to reduce gpu memory usage. + use_der: If True, load and run dereverb model; if False, skip dereverb (default False). Set to False if you want to reduce gpu memory usage. disable_detailed_pbar: Disable detailed progress bars in underlying utils. verbose: Whether to print verbose logs. """ @@ -144,10 +150,15 @@ class VocalSeparator: if verbose: print("[vocal extraction] init: start") - sep_model, sep_config, dereverb_model, dereverb_config, args = build_models(args_dict) + sep_model, sep_config, dereverb_model, dereverb_config, args = build_models(args_dict, use_der=use_der) + sep_model = sep_model.half() sep_model = sep_model.to(device) - dereverb_model = dereverb_model.to(device) + sep_config.inference.chunk_size = int(chunk_length_sec * sep_config.audio.sample_rate) + if dereverb_model is not None: + dereverb_config.inference.chunk_size = int(chunk_length_sec * dereverb_config.audio.sample_rate) + dereverb_model = dereverb_model.half() + dereverb_model = dereverb_model.to(device) self.sep_model = sep_model self.sep_config = sep_config @@ -158,8 +169,9 @@ class VocalSeparator: self.verbose = verbose if verbose: + der_status = "loaded" if dereverb_model is not None else "skipped" print( - "[vocal extraction] init success: sep=loaded, dereverb=loaded, device=", + "[vocal extraction] init success: sep=loaded, dereverb=%s, device=" % der_status, device, ) diff --git a/soulxsinger/models/modules/vocoder.py b/soulxsinger/models/modules/vocoder.py index e739098..bda31f4 100644 --- a/soulxsinger/models/modules/vocoder.py +++ b/soulxsinger/models/modules/vocoder.py @@ -425,8 +425,14 @@ class ISTFTHead(FourierHead): # phase = torch.atan2(y, x) # S = mag * torch.exp(phase * 1j) # better directly produce the complex value + + # Always compute complex values in float32 to avoid ComplexHalf warning (then cast audio back if needed) + orig_dtype = mag.dtype + mag, x, y = mag.float(), x.float(), y.float() S = mag * (x + 1j * y) audio = self.istft(S) + if orig_dtype != torch.float32: + audio = audio.to(orig_dtype) return audio diff --git a/soulxsinger/models/modules/whisper_encoder.py b/soulxsinger/models/modules/whisper_encoder.py new file mode 100644 index 0000000..7b8a0e8 --- /dev/null +++ b/soulxsinger/models/modules/whisper_encoder.py @@ -0,0 +1,74 @@ +"""Frozen Whisper encoder wrapper (wav -> encoder embeddings).""" + +from __future__ import annotations + +from typing import Optional + +import torch +import torch.nn as nn +import torchaudio +from transformers import WhisperFeatureExtractor, WhisperModel + +WHISPER_MEL_FRAMES = 3000 # 3000 frames at 16000 Hz + + +class WhisperEncoder(): + + def __init__( + self, + device: Optional[str] = None, + ) -> None: + self.fe = WhisperFeatureExtractor.from_pretrained("openai/whisper-base") + self.model = WhisperModel.from_pretrained("openai/whisper-base") + self.model = self.model.to(device or ("cuda" if torch.cuda.is_available() else "cpu")) + + def encode( + self, + wav: torch.Tensor, + sr: int, + ) -> torch.Tensor: + wav = torchaudio.functional.resample(wav, orig_freq=sr, new_freq=self.fe.sampling_rate) if sr != self.fe.sampling_rate else wav + wav_np = wav.cpu().detach().numpy().astype("float32", copy=False) + + inputs = self.fe( + wav_np, + sampling_rate=self.fe.sampling_rate, + return_tensors="pt", + padding=False, + truncation=False, + return_attention_mask=True, + ) + + input_features = inputs.input_features + num_frames = input_features.shape[-1] + if num_frames < WHISPER_MEL_FRAMES: + pad = WHISPER_MEL_FRAMES - num_frames + input_features = torch.nn.functional.pad(input_features, (0, pad)) + else: + input_features = input_features[..., :WHISPER_MEL_FRAMES] + + input_features = input_features.to(wav.device) + if self.model.device != wav.device: + self.model = self.model.to(wav.device) + attention_mask = inputs.attention_mask.to(wav.device) if inputs.attention_mask is not None else None + + encoder_out = self.model.encoder(input_features).last_hidden_state + + if attention_mask is not None: + valid_mel_frames = attention_mask.sum(dim=1) + valid_enc_frames = (valid_mel_frames + 1) // 2 + max_valid_enc_frames = min(int(valid_enc_frames.max().item()), encoder_out.shape[1]) + encoder_out = encoder_out[:, :max_valid_enc_frames, :] + valid_len = min(int(valid_enc_frames[0].item()), max_valid_enc_frames) + if valid_len < max_valid_enc_frames: + encoder_out[0, valid_len:, :] = 0 + + return encoder_out + + +if __name__ == "__main__": + torch.manual_seed(0) + audio = torch.randn(1, 24000 * 25).float().to("cuda") + encoder = WhisperEncoder() + whisper_encoder_out = encoder.encode(audio, sr=24000) + print(whisper_encoder_out.shape) diff --git a/soulxsinger/models/soulxsinger.py b/soulxsinger/models/soulxsinger.py index 18aff6c..d4015a8 100644 --- a/soulxsinger/models/soulxsinger.py +++ b/soulxsinger/models/soulxsinger.py @@ -4,6 +4,7 @@ import torch.nn.functional as F import math import numpy as np from typing import Optional, Dict, Any, List +from contextlib import nullcontext from soulxsinger.models.modules.vocoder import Vocoder from soulxsinger.models.modules.decoder import CFMDecoder @@ -11,6 +12,10 @@ from soulxsinger.models.modules.convnext import ConvNeXtV2Block from soulxsinger.models.modules.mel_transform import MelSpectrogramEncoder +def _autocast_if(enabled: bool): + """Return autocast(context) if enabled else no-op context. Use: with _autocast_if(use_amp): ...""" + return torch.amp.autocast(device_type="cuda", enabled=True) if enabled else nullcontext() + class SoulXSinger(nn.Module): """ SoulXSinger model. @@ -102,7 +107,7 @@ class SoulXSinger(nn.Module): return f0_coarse - def infer(self, meta: dict, auto_shift=False, pitch_shift=0, n_steps=32, cfg=3, control="melody"): + def infer(self, meta: dict, auto_shift=False, pitch_shift=0, n_steps=32, cfg=3, control="melody", use_fp16=False): gt_note_text = meta['target']['phoneme'] gt_mel2note = meta['target']['mel2note'] @@ -147,8 +152,13 @@ class SoulXSinger(nn.Module): if gt_note_pitch is None or pt_note_pitch is None: gt_note_pitch, pt_note_pitch = torch.zeros_like(gt_note_type).int().to(gt_note_type.device), torch.zeros_like(pt_note_type).int().to(pt_note_type.device) - # convert prompt waveform to mel spectrogram - pt_mel = self.mel(pt_wav) + use_fp16 = use_fp16 and pt_wav.is_cuda + # mel is kept in fp32 (see build_model: model.mel.float() after model.half()) + pt_mel = self.mel(pt_wav.float() if pt_wav.dtype != torch.float32 else pt_wav) + if use_fp16: + pt_mel = pt_mel.half() + pt_f0 = pt_f0.half() + gt_f0 = gt_f0.half() len_prompt = pt_note_pitch.shape[1] len_prompt_mel = pt_f0.shape[1] @@ -165,23 +175,23 @@ class SoulXSinger(nn.Module): note_pitch[note_pitch > 0] = note_pitch[note_pitch > 0] + f0_shift note_pitch = torch.clamp(note_pitch, 0, 255) - features = self.note_pitch_encoder(note_pitch) + self.note_type_encoder(note_type) + self.note_text_encoder(note_text) - - features = self.preflow(features) - features = self.expand_states(features, mel2note) - features = features + self.f0_encoder(f0_course) - - gt_decoder_inp = features[:, len_prompt_mel:, :] - pt_decoder_inp = features[:, :len_prompt_mel, :] + with _autocast_if(use_fp16): + features = self.note_pitch_encoder(note_pitch) + self.note_type_encoder(note_type) + self.note_text_encoder(note_text) + + features = self.preflow(features) + features = self.expand_states(features, mel2note) + features = features + self.f0_encoder(f0_course) + + gt_decoder_inp = features[:, len_prompt_mel:, :] + pt_decoder_inp = features[:, :len_prompt_mel, :] + + generated_mel = self.cfm_decoder.reverse_diffusion( + pt_mel, + pt_decoder_inp, + gt_decoder_inp, + n_timesteps=n_steps, + cfg=cfg + ) + generated_audio = self.vocoder(generated_mel.transpose(1, 2)[0:1, ...]).float() - generated_mel = self.cfm_decoder.reverse_diffusion( - pt_mel, - pt_decoder_inp, - gt_decoder_inp, - n_timesteps=n_steps, - cfg=cfg - ) - - generated_audio = self.vocoder(generated_mel.transpose(1, 2)[0:1, ...]) - return generated_audio diff --git a/soulxsinger/models/soulxsinger_svc.py b/soulxsinger/models/soulxsinger_svc.py new file mode 100644 index 0000000..1cfbe62 --- /dev/null +++ b/soulxsinger/models/soulxsinger_svc.py @@ -0,0 +1,339 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F +import numpy as np +from tqdm import tqdm +from typing import Optional, Dict, Any, List, Tuple +from contextlib import nullcontext + +from soulxsinger.models.modules.vocoder import Vocoder +from soulxsinger.models.modules.decoder import CFMDecoder +from soulxsinger.models.modules.mel_transform import MelSpectrogramEncoder +from soulxsinger.models.modules.whisper_encoder import WhisperEncoder + + +def _autocast_if(enabled: bool): + """Return autocast(context) if enabled else no-op context. Use: with _autocast_if(use_amp): ...""" + return torch.amp.autocast(device_type="cuda", enabled=True) if enabled else nullcontext() + +class SoulXSingerSVC(nn.Module): + """ + SoulXSinger SVC model. + """ + def __init__(self, config: Dict): + super(SoulXSingerSVC, self).__init__() + self.audio_cfg = config.audio + enc_cfg = config.model.encoder + cfm_cfg = config.model.flow_matching + + self.whisper_encoder = WhisperEncoder() + self.f0_encoder = nn.Embedding(enc_cfg["f0_bin"], enc_cfg["f0_dim"]) + self.cfm_decoder = CFMDecoder(cfm_cfg) + + self.mel = MelSpectrogramEncoder(self.audio_cfg) + self.vocoder = Vocoder() + + @staticmethod + def f0_to_coarse(f0, f0_bin=361, f0_min=32.7031956625, f0_shift=0): + """ + Convert continuous F0 values to discrete F0 bins (SIL and C1 - B6, 361 bins). + args: + f0: continuous F0 values + f0_bin: number of F0 bins + f0_min: minimum F0 value + f0_shift: shift value for F0 bins + returns: + f0_coarse: discrete F0 bins + """ + is_torch = isinstance(f0, torch.Tensor) + uv_mask = f0 <= 0 + + if is_torch: + f0_safe = torch.maximum(f0, torch.tensor(f0_min)) + f0_cents = 1200 * torch.log2(f0_safe / f0_min) + else: + f0_safe = np.maximum(f0, f0_min) + f0_cents = 1200 * np.log2(f0_safe / f0_min) + + f0_coarse = (f0_cents / 20) + 1 + + if is_torch: + f0_coarse = torch.round(f0_coarse).long() + f0_coarse = torch.clamp(f0_coarse, min=1, max=f0_bin - 1) + else: + f0_coarse = np.rint(f0_coarse).astype(int) + f0_coarse = np.clip(f0_coarse, 1, f0_bin - 1) + + f0_coarse[uv_mask] = 0 + + if f0_shift != 0: + if is_torch: + voiced = f0_coarse > 0 + if voiced.any(): + shifted = f0_coarse[voiced] + f0_shift + f0_coarse[voiced] = torch.clamp(shifted, 1, f0_bin - 1) + else: + voiced = f0_coarse > 0 + if np.any(voiced): + shifted = f0_coarse[voiced] + f0_shift + f0_coarse[voiced] = np.clip(shifted, 1, f0_bin - 1) + + return f0_coarse + + @staticmethod + def build_vocal_segments( + f0, + f0_rate: int = 50, + uv_frames_th: int = 5, + min_duration_sec: float = 5.0, + max_duration_sec: float = 30.0, + num_overlaps: int = 1, + ignore_silent_segments: bool = True, + ) -> Tuple[List[Tuple[float, float]], List[Tuple[float, float]]]: + """Build vocal segments based on F0 contour. First split by long silent runs, then merge into segments based on min and max duration constraints. + args: + f0: F0 contour of the audio, 1D array or tensor with shape (T,) + f0_rate: F0 sampling rate in Hz (e.g., 50 for 20ms hop size) + uv_frames_th: number of consecutive zero F0 frames to consider as a split point + min_duration_sec: minimum duration of each segment in seconds + max_duration_sec: maximum duration of each segment in seconds + num_overlaps: number of overlapping segments to create for each non-overlapping segment (for smooth inference) + ignore_silent_segments: whether to ignore segments that are mostly silent (e.g., > 95% zero F0) + returns: + overlap_segments: list of (overlap_start_sec, overlap_end_sec) for each segment, which may overlap with adjacent segments for smooth inference + segments: list of (seg_start_sec, seg_end_sec) for each segment, which are non-overlapping and used for final merging + """ + if isinstance(f0, torch.Tensor): + f0_np = f0.detach().float().cpu().numpy() + else: + f0_np = np.asarray(f0, dtype=np.float32) + f0_np = np.squeeze(f0_np) + + total_frames = int(f0_np.shape[0]) + if total_frames == 0: + return [], [] + + min_frames = max(1, int(round(min_duration_sec * f0_rate))) + max_frames = max(1, int(round(max_duration_sec * f0_rate))) + + split_points = [0] # silence split points in frame indices, starting with 0 and ending with total_frames + + def append_split_point(point: int): + # Ensure split points are within valid range and respect max_frames constraint + point = int(max(0, min(point, total_frames))) + while point - split_points[-1] > max_frames: + split_points.append(split_points[-1] + max_frames) + if point > split_points[-1]: + split_points.append(point) + + idx = 0 + while idx < total_frames: + if f0_np[idx] == 0: + run_start = idx + while idx < total_frames and f0_np[idx] == 0: + idx += 1 + run_end = idx + if (run_end - run_start) >= uv_frames_th: + split_point = max(run_end - 5, (run_start + run_end) // 2) + append_split_point(split_point) + else: + idx += 1 + append_split_point(total_frames) + # print(f"Initial split points (in seconds): {[round(p / f0_rate, 2) for p in split_points]}") + + segments: List[Tuple[int, int]] = [] + overlap_segments: List[Tuple[int, int]] = [] + + def append_segment(start_idx: int, end_idx: int, num_overlaps: int = num_overlaps): + segments.append((split_points[start_idx] / f0_rate, split_points[end_idx] / f0_rate)) + overlap_start_idx = start_idx + if start_idx > 0 and (split_points[end_idx] - split_points[start_idx - num_overlaps]) <= max_frames: + overlap_start_idx = start_idx - num_overlaps + overlap_segments.append((split_points[overlap_start_idx] / f0_rate, split_points[end_idx] / f0_rate)) + + segment_start, segment_end = 0, 1 + + while segment_start < len(split_points) - 1: + while segment_end < len(split_points) and (split_points[segment_end] - split_points[segment_start]) < min_frames: + segment_end += 1 + + if segment_end >= len(split_points): + append_segment(segment_start, len(split_points) - 1, num_overlaps=num_overlaps) + break + append_segment(segment_start, segment_end, num_overlaps=num_overlaps) + segment_start = segment_end + segment_end = segment_start + 1 + + # print(f"Final segments (overlap_start, overlap_end, seg_start_time, seg_end_time) in seconds: {overlap_segments}") + if ignore_silent_segments: + filtered_idx = [] + for i, seg in enumerate(overlap_segments): + start_frame = int(seg[0] * f0_rate) + end_frame = int(seg[1] * f0_rate) + total_frames = end_frame - start_frame + voice_frames = np.sum(f0_np[start_frame:end_frame] > 0) + if voice_frames / total_frames > 0.05 and voice_frames >= 10: # at least 10 voiced frames and >5% voiced frames + filtered_idx.append(i) + + overlap_segments = [overlap_segments[i] for i in filtered_idx] + segments = [segments[i] for i in filtered_idx] + # print(f"Filtered segments with mostly silence removed: {overlap_segments}") + + return overlap_segments, segments + + def infer( + self, + pt_wav: str|torch.Tensor, + gt_wav: str|torch.Tensor, + pt_f0: str|torch.Tensor, + gt_f0: str|torch.Tensor, + auto_shift=False, + pitch_shift=0, + n_steps=32, + cfg=3, + use_fp16=False, + ): + """ + SVC inference pipeline. First build vocal segments based on F0 contour, then run inference for each segment and merge results. + args: + pt_wav: prompt waveform path or tensor + gt_wav: target waveform path or tensor + pt_f0: prompt F0 path or tensor + gt_f0: target F0 path or tensor + auto_shift: whether to automatically calculate pitch shift based on median F0 of prompt and target + pitch_shift: manual pitch shift in semitones (overrides auto_shift if > 0) + n_steps: number of diffusion steps for inference + cfg: classifier-free guidance scale for inference + use_fp16: if True, run in FP16 except mel extraction to save memory and speed. + """ + + # calculate auto pitch shift + if auto_shift and pitch_shift == 0: + if gt_f0 is not None and pt_f0 is not None: + gt_f0_median = torch.median(gt_f0[gt_f0 > 0]) + pt_f0_median = torch.median(pt_f0[pt_f0 > 0]) + pitch_shift = torch.round(torch.log2(pt_f0_median / gt_f0_median) * 1200 / 100).int().item() + else: + print("Warning: pitch_shift is True but note_pitch or f0 is None. Set f0_shift to 0.") + pitch_shift = 0 + else: + pitch_shift = pitch_shift + + use_fp16 = use_fp16 and pt_wav.is_cuda + # mel is kept in fp32 (see build_model: model.mel.float() after model.half()) + pt_mel = self.mel(pt_wav.float() if pt_wav.dtype != torch.float32 else pt_wav) + if use_fp16: + pt_mel = pt_mel.half() + pt_wav = pt_wav.half() + gt_wav = gt_wav.half() + pt_f0 = pt_f0.half() + gt_f0 = gt_f0.half() + + # if target audio is less than 30 seconds, infer the whole audio + if gt_wav.shape[-1] < 30 * self.audio_cfg.sample_rate: + with _autocast_if(use_fp16): + generated_audio = self.infer_segment( + pt_mel=pt_mel, + pt_wav=pt_wav, + gt_wav=gt_wav, + pt_f0=pt_f0, + gt_f0=gt_f0, + pitch_shift=pitch_shift, + n_steps=n_steps, + cfg=cfg, + ) + return generated_audio, pitch_shift + + # if target audio is longer than 30 seconds, build vocal segments and infer each segment + generated_audio = [] + + f0_rate = self.audio_cfg.sample_rate // self.audio_cfg.hop_size + + overlap_segments, segments = self.build_vocal_segments( + gt_f0, + f0_rate=f0_rate, + uv_frames_th=10, + min_duration_sec=15.0, + max_duration_sec=30.0, + ) + if len(segments) == 0: + segments = [(0.0, gt_wav.shape[-1] / self.audio_cfg.sample_rate)] + overlap_segments = [(0.0, gt_wav.shape[-1] / self.audio_cfg.sample_rate)] + + generated_audio = torch.zeros_like(gt_wav) + for idx in tqdm(range(len(segments)), total=len(segments), desc="Inferring segments (SVC)", dynamic_ncols=True): + overlap_start_sec, overlap_end_sec = overlap_segments[idx] + seg_start_sec, seg_end_sec = segments[idx] + + wav_start = int(round(overlap_start_sec * self.audio_cfg.sample_rate)) + wav_end = int(round(overlap_end_sec * self.audio_cfg.sample_rate)) + f0_start = int(round(overlap_start_sec * f0_rate)) + f0_end = int(round(overlap_end_sec * f0_rate)) + + wav_start = max(0, min(wav_start, gt_wav.shape[-1])) + wav_end = max(wav_start, min(wav_end, gt_wav.shape[-1])) + f0_start = max(0, min(f0_start, gt_f0.shape[-1])) + f0_end = max(f0_start, min(f0_end, gt_f0.shape[-1])) + + segment_gt_wav = gt_wav[:, wav_start:wav_end] + segment_gt_f0 = gt_f0[:, f0_start:f0_end] + with _autocast_if(use_fp16): + segment_generated_audio = self.infer_segment( + pt_mel=pt_mel, + pt_wav=pt_wav, + gt_wav=segment_gt_wav, + pt_f0=pt_f0, + gt_f0=segment_gt_f0, + pitch_shift=pitch_shift, + n_steps=n_steps, + cfg=cfg, + ) + + segment_start = int(round(seg_start_sec * self.audio_cfg.sample_rate)) + segment_end = int(round(seg_end_sec * self.audio_cfg.sample_rate)) + segment_generated_audio = segment_generated_audio[segment_start - wav_start: segment_end - wav_start] + + generated_audio[:, segment_start:segment_end] = segment_generated_audio + + return generated_audio, pitch_shift + + def infer_segment(self, pt_mel, pt_wav, gt_wav, pt_f0, gt_f0, pitch_shift=0, n_steps=32, cfg=3): + len_prompt_mel = pt_mel.shape[1] + pt_f0 = F.pad(pt_f0, (0, 0, 0, max(0, len_prompt_mel - pt_f0.shape[1])))[:, :len_prompt_mel] + + f0_course_pt = self.f0_to_coarse(pt_f0) + f0_course_gt = self.f0_to_coarse(gt_f0, f0_shift=pitch_shift * 5) + f0_course = torch.cat([f0_course_pt, f0_course_gt], 1) + + pt_content_feat = self.whisper_encoder.encode(pt_wav, sr=self.audio_cfg.sample_rate) + gt_content_feat = self.whisper_encoder.encode(gt_wav, sr=self.audio_cfg.sample_rate) + t_pt, t_gt = f0_course_pt.shape[1], f0_course_gt.shape[1] + pt_content_feat = F.pad(pt_content_feat, (0, 0, 0, max(0, t_pt - pt_content_feat.shape[1])))[:, :t_pt, :] + gt_content_feat = F.pad(gt_content_feat, (0, 0, 0, max(0, t_gt - gt_content_feat.shape[1])))[:, :t_gt, :] + + content_feat = torch.cat([pt_content_feat, gt_content_feat], 1) + + f0_feat = self.f0_encoder(f0_course) + features = content_feat + f0_feat + + gt_decoder_inp = features[:, len_prompt_mel:, :] + pt_decoder_inp = features[:, :len_prompt_mel, :] + + generated_mel = self.cfm_decoder.reverse_diffusion( + pt_mel, + pt_decoder_inp, + gt_decoder_inp, + n_timesteps=n_steps, + cfg=cfg + ) + + generated_audio = self.vocoder(generated_mel.transpose(1, 2)[0:1, ...]) + generated_audio = generated_audio.squeeze().float() + + # cut or pad to match gt_wav length + if generated_audio.shape[-1] > gt_wav.shape[-1]: + generated_audio = generated_audio[:gt_wav.shape[-1]] + elif generated_audio.shape[-1] < gt_wav.shape[-1]: + generated_audio = F.pad(generated_audio, (0, gt_wav.shape[-1] - generated_audio.shape[-1])) + + return generated_audio \ No newline at end of file diff --git a/webui.py b/webui.py index 4b3a4ee..48c8cfc 100644 --- a/webui.py +++ b/webui.py @@ -273,8 +273,9 @@ def _control_to_internal(control: str) -> str: class AppState: - def __init__(self) -> None: + def __init__(self, use_fp16: bool = False) -> None: self.device = _get_device() + self.use_fp16 = use_fp16 and ("cuda" in self.device) self.preprocess_pipeline = PreprocessPipeline( device=self.device, language="Mandarin", @@ -288,6 +289,7 @@ class AppState: model_path="pretrained_models/SoulX-Singer/model.pt", config=config, device=self.device, + use_fp16=self.use_fp16, ) self.phoneset_path = "soulxsinger/utils/phoneme/phone_set.json" self.midi_parser = MidiParser( @@ -345,6 +347,7 @@ class AppState: args.auto_shift = auto_shift args.pitch_shift = int(pitch_shift) args.control = control + args.use_fp16 = self.use_fp16 try: svs_process(args, self.svs_config, self.svs_model) gc.collect() @@ -392,7 +395,7 @@ class AppState: return True, "svs inference done", merged -APP_STATE = AppState() +APP_STATE = AppState(use_fp16="--fp16" in sys.argv) def _edit_metadata( meta, @@ -880,6 +883,7 @@ if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("--port", type=int, default=7860, help="Gradio server port") parser.add_argument("--share", action="store_true", help="Create public link") + parser.add_argument("--fp16", action="store_true", help="Use FP16 for SVS model and inference") args = parser.parse_args() page = render_interface() diff --git a/webui_svc.py b/webui_svc.py new file mode 100644 index 0000000..52fe0c4 --- /dev/null +++ b/webui_svc.py @@ -0,0 +1,465 @@ +import random +import sys +import traceback +import gc +from datetime import datetime +from pathlib import Path +from typing import Literal + +import gradio as gr +import librosa +import numpy as np +import soundfile as sf +import torch + +from preprocess.pipeline import PreprocessPipeline +from soulxsinger.utils.file_utils import load_config +from cli.inference_svc import build_model as build_svc_model, process as svc_process + + +ROOT = Path(__file__).parent +SAMPLE_RATE = 44100 +PROMPT_MAX_SEC_DEFAULT = 30 +TARGET_MAX_SEC_DEFAULT = 600 + +SVC_EXAMPLE_PROMPT_AUDIO = "example/audio/svc_prompt_demo.mp3" +SVC_EXAMPLE_TARGET_AUDIO = "example/audio/svc_target_demo.mp3" + +EXAMPLE_LIST = [[ + str(ROOT / SVC_EXAMPLE_PROMPT_AUDIO), + str(ROOT / SVC_EXAMPLE_TARGET_AUDIO), + False, + True, + True, + True, + 0, + 32, + 1.0, + 42, +]] + +_I18N = dict( + display_lang_label=dict(en="Display Language", zh="显示语言"), + title=dict(en="## SoulX-Singer SVC", zh="## SoulX-Singer SVC"), + prompt_audio_label=dict(en=f"Prompt audio", zh=f"Prompt 音频"), + target_audio_label=dict(en=f"Target audio", zh=f"Target 音频"), + prompt_vocal_sep_label=dict(en="Prompt vocal separation", zh="Prompt 人声分离"), + target_vocal_sep_label=dict(en="Target vocal separation", zh="Target 人声分离"), + auto_shift_label=dict(en="Auto pitch shift", zh="自动变调"), + auto_mix_acc_label=dict(en="Auto mix accompaniment", zh="自动混合伴奏"), + pitch_shift_label=dict(en="Pitch shift (semitones)", zh="指定变调(半音)"), + n_step_label=dict(en="n_step", zh="采样步数"), + cfg_label=dict(en="cfg scale", zh="cfg系数"), + seed_label=dict(en="Seed", zh="种子"), + examples_label=dict(en="Reference example (click to load)", zh="参考样例(点击加载)"), + run_btn=dict(en="🎤Singing Voice Conversion", zh="🎤歌声转换"), + output_audio_label=dict(en="Generated audio", zh="合成结果音频"), + warn_missing_audio=dict(en="Please provide both prompt audio and target audio.", zh="请同时上传 Prompt 与 Target 音频。"), + instruction_title=dict(en="Usage", zh="使用说明"), + instruction_p1=dict( + en="Upload the Prompt and Target audio, and configure the parameters", + zh="上传 Prompt 与 Target 音频,并配置相关参数", + ), + instruction_p2=dict( + en="Click「🎤Singing Voice Conversion」to start singing voice conversion.", + zh="点击「🎤歌声转换」开始最终生成。", + ), + tips_title=dict(en="Tips", zh="提示"), + tip_p1=dict( + en="Input: The Prompt audio is recommended to be a clean and clear singing voice, while the Target audio can be either a pure vocal or a mixture with accompaniment. If the audio contains accompaniment, please check the vocal separation option.", + zh="输入:Prompt 音频建议是干净清晰的歌声,Target 音频可以是纯歌声或伴奏,这两者若带伴奏需要勾选分离选项", + ), + tip_p2=dict( + en="Pitch shift: When there is a large pitch range difference between the Prompt and Target audio, you can try enabling auto pitch shift or manually adjusting the pitch shift in semitones. When a non-zero pitch shift is specified, auto pitch shift will not take effect. The accompaniment of auto mix will be pitch-shifted together with the vocal (keeping the same octave).", + zh="变调:Prompt 音频的音域和 Target 音频的音域差距较大的时候,可以尝试开启自动变调或手动调整变调半音数,指定非0的变调半音数时,自动变调不生效,自动混音的伴奏会配合歌声进行升降调(保持同一个八度)", + ), + tip_p3=dict( + en="Model parameters: Generally, a larger number of sampling steps will yield better generation quality but also longer generation time; a larger cfg scale will increase timbre similarity and melody fidelity, but may cause more distortion, it is recommended to take a value between 1 and 3.", + zh="模型参数:一般采样步数越大,生成质量越好,但生成时间也越长;一般cfg系数越大,音色相似度和旋律保真度越高,但是会造成更多的失真,建议取1~3之间的值", + ), + tip_p4=dict( + en="If you want to convert a long audio or a whole song with large pitch range, there may be instability in the generated voice. You can try converting in segments.", + zh="长音频或完整歌曲中,音域变化较大的情况有可能出现音色不稳定,可以尝试分段转换", + ) +) + +_GLOBAL_LANG: Literal["zh", "en"] = "zh" + + +def _i18n(key: str) -> str: + return _I18N[key][_GLOBAL_LANG] + + +def _print_exception(context: str) -> None: + print(f"[{context}]\n{traceback.format_exc()}", file=sys.stderr, flush=True) + + +def _get_device() -> str: + return "cuda:0" if torch.cuda.is_available() else "cpu" + + +def _session_dir() -> Path: + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S_%f") + return ROOT / "outputs" / "gradio" / "svc" / timestamp + + +def _normalize_audio_input(audio): + return audio[0] if isinstance(audio, tuple) else audio + + +def _trim_and_save_audio(src_audio_path: str, dst_wav_path: Path, max_sec: int, sr: int = SAMPLE_RATE) -> None: + audio_data, _ = librosa.load(src_audio_path, sr=sr, mono=True) + audio_data = audio_data[: max_sec * sr] + dst_wav_path.parent.mkdir(parents=True, exist_ok=True) + sf.write(dst_wav_path, audio_data, sr) + + +def _usage_md() -> str: + return "\n\n".join([ + f"### {_i18n('instruction_title')}", + f"**1.** {_i18n('instruction_p1')}", + f"**2.** {_i18n('instruction_p2')}", + ]) + + +def _tips_md() -> str: + return "\n\n".join([ + f"### {_i18n('tips_title')}", + f"- {_i18n('tip_p1')}", + f"- {_i18n('tip_p2')}", + f"- {_i18n('tip_p3')}", + f"- {_i18n('tip_p4')}", + ]) + + +class AppState: + def __init__(self, use_fp16: bool = False) -> None: + self.device = _get_device() + self.use_fp16 = use_fp16 and ("cuda" in self.device) + self.preprocess_pipeline = PreprocessPipeline( + device=self.device, + language="Mandarin", + save_dir=str(ROOT / "outputs" / "gradio" / "_placeholder" / "svc"), + vocal_sep=True, + max_merge_duration=60000, + midi_transcribe=False, + ) + + self.svc_config = load_config("soulxsinger/config/soulxsinger.yaml") + self.svc_model = build_svc_model( + model_path="pretrained_models/SoulX-Singer/model-svc.pt", + config=self.svc_config, + device=self.device, + use_fp16=self.use_fp16, + ) + + def run_preprocess(self, audio_path: Path, save_path: Path, vocal_sep: bool) -> tuple[bool, str, Path | None, Path | None]: + try: + self.preprocess_pipeline.save_dir = str(save_path) + self.preprocess_pipeline.run( + audio_path=str(audio_path), + vocal_sep=vocal_sep, + max_merge_duration=60000, + language="Mandarin", + ) + vocal_wav = save_path / "vocal.wav" + vocal_f0 = save_path / "vocal_f0.npy" + if not vocal_wav.exists() or not vocal_f0.exists(): + return False, f"preprocess output missing: {vocal_wav} or {vocal_f0}", None, None + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + return True, "ok", vocal_wav, vocal_f0 + except Exception as e: + return False, f"preprocess failed: {e}", None, None + + def run_svc( + self, + prompt_wav_path: Path, + target_wav_path: Path, + prompt_f0_path: Path, + target_f0_path: Path, + session_base: Path, + auto_shift: bool, + auto_mix_acc: bool, + pitch_shift: int, + n_step: int, + cfg: float, + seed: int, + ) -> tuple[bool, str, Path | None]: + try: + torch.manual_seed(seed) + np.random.seed(seed) + random.seed(seed) + + save_dir = session_base / "generated" + save_dir.mkdir(parents=True, exist_ok=True) + + class Args: + pass + + args = Args() + args.device = self.device + args.prompt_wav_path = str(prompt_wav_path) + args.target_wav_path = str(target_wav_path) + args.prompt_f0_path = str(prompt_f0_path) + args.target_f0_path = str(target_f0_path) + args.save_dir = str(save_dir) + args.auto_shift = auto_shift + args.pitch_shift = int(pitch_shift) + args.n_steps = int(n_step) + args.cfg = float(cfg) + args.use_fp16 = self.use_fp16 + + svc_process(args, self.svc_config, self.svc_model) + + generated = save_dir / "generated.wav" + if not generated.exists(): + return False, f"inference finished but output not found: {generated}", None + + if auto_mix_acc: + acc_path = session_base / "transcriptions" / "target" / "acc.wav" + if acc_path.exists(): + vocal_shift = args.pitch_shift + mul = -1 if vocal_shift < 0 else 1 + acc_shift = abs(vocal_shift) % 12 + acc_shift = mul * acc_shift + if acc_shift > 6: + acc_shift -= 12 + if acc_shift < -6: + acc_shift += 12 + + mix_sr = self.svc_config.audio.sample_rate + vocal, _ = librosa.load(str(generated), sr=mix_sr, mono=True) + acc, _ = librosa.load(str(acc_path), sr=mix_sr, mono=True) + if acc_shift != 0: + acc = librosa.effects.pitch_shift(acc, sr=mix_sr, n_steps=acc_shift) + print(f"Applied pitch shift of {acc_shift} semitones to accompaniment to match vocal shift of {vocal_shift} semitones.") + + mix_len = min(len(vocal), len(acc)) + if mix_len > 0: + mixed = vocal[:mix_len] + acc[:mix_len] + peak = float(np.max(np.abs(mixed))) if mixed.size > 0 else 1.0 + if peak > 1.0: + mixed = mixed / peak + mixed_path = save_dir / "generated_mixed.wav" + sf.write(str(mixed_path), mixed, mix_sr) + generated = mixed_path + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + return True, "svc inference done", generated + except Exception as e: + return False, f"svc inference failed: {e}", None + + +APP_STATE = AppState(use_fp16="--fp16" in sys.argv) + + +def _start_svc(prompt_audio, target_audio, prompt_vocal_sep, target_vocal_sep, auto_shift, auto_mix_acc, pitch_shift, n_step, cfg, seed): + try: + prompt_audio = _normalize_audio_input(prompt_audio) + target_audio = _normalize_audio_input(target_audio) + if not prompt_audio or not target_audio: + gr.Warning(_i18n("warn_missing_audio")) + return None + + session_base = _session_dir() + audio_dir = session_base / "audio" + prompt_raw = audio_dir / "prompt.wav" + target_raw = audio_dir / "target.wav" + _trim_and_save_audio(prompt_audio, prompt_raw, PROMPT_MAX_SEC_DEFAULT) + _trim_and_save_audio(target_audio, target_raw, TARGET_MAX_SEC_DEFAULT) + + prompt_ok, prompt_msg, prompt_wav, prompt_f0 = APP_STATE.run_preprocess( + audio_path=prompt_raw, + save_path=session_base / "transcriptions" / "prompt", + vocal_sep=bool(prompt_vocal_sep), + ) + if not prompt_ok or prompt_wav is None or prompt_f0 is None: + print(prompt_msg, file=sys.stderr, flush=True) + return None + + target_ok, target_msg, target_wav, target_f0 = APP_STATE.run_preprocess( + audio_path=target_raw, + save_path=session_base / "transcriptions" / "target", + vocal_sep=bool(target_vocal_sep), + ) + if not target_ok or target_wav is None or target_f0 is None: + print(target_msg, file=sys.stderr, flush=True) + return None + + ok, msg, generated = APP_STATE.run_svc( + prompt_wav_path=prompt_wav, + target_wav_path=target_wav, + prompt_f0_path=prompt_f0, + target_f0_path=target_f0, + session_base=session_base, + auto_shift=bool(auto_shift), + auto_mix_acc=bool(auto_mix_acc), + pitch_shift=int(pitch_shift), + n_step=int(n_step), + cfg=float(cfg), + seed=int(seed), + ) + if not ok or generated is None: + print(msg, file=sys.stderr, flush=True) + return None + return str(generated) + except Exception: + _print_exception("_start_svc") + return None + + +def render_interface() -> gr.Blocks: + with gr.Blocks(title="SoulX-Singer-SVC Demo", theme=gr.themes.Default()) as page: + gr.HTML( + '
' + '
SoulX-Singer-SVC
' + '
' + '
' + ) + with gr.Row(equal_height=True): + lang_choice = gr.Radio( + choices=["中文", "English"], + value="中文", + label=_i18n("display_lang_label"), + type="index", + interactive=True, + ) + + usage_md = gr.Markdown(_usage_md()) + + with gr.Row(equal_height=True): + prompt_audio = gr.Audio( + label=_i18n("prompt_audio_label"), + type="filepath", + editable=False, + interactive=True, + ) + target_audio = gr.Audio( + label=_i18n("target_audio_label"), + type="filepath", + editable=False, + interactive=True, + ) + + with gr.Row(equal_height=True): + prompt_vocal_sep = gr.Checkbox(label=_i18n("prompt_vocal_sep_label"), value=False, scale=1) + target_vocal_sep = gr.Checkbox(label=_i18n("target_vocal_sep_label"), value=True, scale=1) + auto_shift = gr.Checkbox(label=_i18n("auto_shift_label"), value=True, scale=1) + auto_mix_acc = gr.Checkbox(label=_i18n("auto_mix_acc_label"), value=True, scale=1) + + with gr.Row(equal_height=True): + pitch_shift = gr.Slider(label=_i18n("pitch_shift_label"), value=0, minimum=-36, maximum=36, step=1, scale=1) + n_step = gr.Slider(label=_i18n("n_step_label"), value=32, minimum=1, maximum=200, step=1, scale=1) + cfg = gr.Slider(label=_i18n("cfg_label"), value=1.0, minimum=0.0, maximum=10.0, step=0.1, scale=1) + seed_input = gr.Slider(label=_i18n("seed_label"), value=42, minimum=0, maximum=10000, step=1, scale=1) + + with gr.Row(): + run_btn = gr.Button(value=_i18n("run_btn"), variant="primary", size="lg") + + with gr.Row(): + output_audio = gr.Audio(label=_i18n("output_audio_label"), type="filepath", interactive=False) + + gr.Examples( + examples=EXAMPLE_LIST, + inputs=[prompt_audio, target_audio], + label=_i18n("examples_label"), + ) + + tips_md = gr.Markdown(_tips_md()) + + run_btn.click( + fn=_start_svc, + inputs=[ + prompt_audio, + target_audio, + prompt_vocal_sep, + target_vocal_sep, + auto_shift, + auto_mix_acc, + pitch_shift, + n_step, + cfg, + seed_input, + ], + outputs=[output_audio], + ) + + def _change_language(lang): + global _GLOBAL_LANG + _GLOBAL_LANG = ["zh", "en"][lang] + return [ + gr.update(label=_i18n("display_lang_label")), + gr.update(value=_i18n("title")), + gr.update(value=_usage_md()), + gr.update(label=_i18n("prompt_audio_label")), + gr.update(label=_i18n("target_audio_label")), + gr.update(label=_i18n("prompt_vocal_sep_label")), + gr.update(label=_i18n("target_vocal_sep_label")), + gr.update(label=_i18n("auto_shift_label")), + gr.update(label=_i18n("auto_mix_acc_label")), + gr.update(label=_i18n("pitch_shift_label")), + gr.update(label=_i18n("n_step_label")), + gr.update(label=_i18n("cfg_label")), + gr.update(label=_i18n("seed_label")), + gr.update(value=_i18n("run_btn")), + gr.update(label=_i18n("output_audio_label")), + gr.update(value=_tips_md()), + ] + + lang_choice.change( + fn=_change_language, + inputs=[lang_choice], + outputs=[ + lang_choice, + usage_md, + prompt_audio, + target_audio, + prompt_vocal_sep, + target_vocal_sep, + auto_shift, + auto_mix_acc, + pitch_shift, + n_step, + cfg, + seed_input, + run_btn, + output_audio, + tips_md, + ], + ) + + return page + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser() + parser.add_argument("--port", type=int, default=7861, help="Gradio server port") + parser.add_argument("--share", action="store_true", help="Create public link") + parser.add_argument("--fp16", action="store_true", help="Use FP16 for SVC model and inference") + args = parser.parse_args() + + page = render_interface() + page.queue() + page.launch(share=args.share, server_name="0.0.0.0", server_port=args.port)