From 6bc423f3af5ad12c9425ee0a24d5039a80f2c83c Mon Sep 17 00:00:00 2001 From: jlqian98 Date: Mon, 16 Mar 2026 13:40:30 +0800 Subject: [PATCH] add fp16 support for svs and svc inference --- cli/inference.py | 21 +++++++-- cli/inference_svc.py | 37 +++++++++++++--- example/infer.sh | 3 +- example/infer_svc.sh | 3 +- soulxsinger/models/modules/vocoder.py | 6 +++ soulxsinger/models/soulxsinger.py | 52 +++++++++++++--------- soulxsinger/models/soulxsinger_svc.py | 62 ++++++++++++++++++--------- webui.py | 8 +++- webui_svc.py | 8 +++- 9 files changed, 143 insertions(+), 57 deletions(-) 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 index e2c335f..8372b50 100644 --- a/cli/inference_svc.py +++ b/cli/inference_svc.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. + SoulXSingerSVC: 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,9 +49,13 @@ 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).") + print("Model checkpoint loaded.") model.eval() model.to(device) - print("Model checkpoint loaded.") return model @@ -67,8 +73,19 @@ def process(args, config, model: torch.nn.Module): 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 - generated_audio, generated_shift = model.infer(pt_wav, gt_wav, pt_f0, gt_f0, auto_shift=args.auto_shift, pitch_shift=args.pitch_shift, n_steps=n_step, cfg=cfg) - generated_audio = generated_audio.squeeze().cpu().numpy() + 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.") @@ -82,6 +99,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) @@ -99,7 +117,14 @@ if __name__ == "__main__": 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/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 index 8fbc484..036ef83 100644 --- a/example/infer_svc.sh +++ b/example/infer_svc.sh @@ -24,4 +24,5 @@ python -m cli.inference_svc \ --target_f0_path $target_f0_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/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/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 index d04fbb9..1cfbe62 100644 --- a/soulxsinger/models/soulxsinger_svc.py +++ b/soulxsinger/models/soulxsinger_svc.py @@ -4,6 +4,7 @@ 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 @@ -11,6 +12,10 @@ 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. @@ -186,6 +191,7 @@ class SoulXSingerSVC(nn.Module): 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. @@ -198,6 +204,7 @@ class SoulXSingerSVC(nn.Module): 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 @@ -212,17 +219,29 @@ class SoulXSingerSVC(nn.Module): 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: - generated_audio = self.infer_segment( - 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, - ) + 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 @@ -258,15 +277,17 @@ class SoulXSingerSVC(nn.Module): segment_gt_wav = gt_wav[:, wav_start:wav_end] segment_gt_f0 = gt_f0[:, f0_start:f0_end] - segment_generated_audio = self.infer_segment( - 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, - ) + 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)) @@ -276,8 +297,7 @@ class SoulXSingerSVC(nn.Module): return generated_audio, pitch_shift - def infer_segment(self, pt_wav, gt_wav, pt_f0, gt_f0, pitch_shift=0, n_steps=32, cfg=3): - pt_mel = self.mel(pt_wav) + 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] @@ -308,7 +328,7 @@ class SoulXSingerSVC(nn.Module): ) generated_audio = self.vocoder(generated_mel.transpose(1, 2)[0:1, ...]) - generated_audio = generated_audio.squeeze() + generated_audio = generated_audio.squeeze().float() # cut or pad to match gt_wav length if generated_audio.shape[-1] > gt_wav.shape[-1]: diff --git a/webui.py b/webui.py index ac9db6b..57918e4 100644 --- a/webui.py +++ b/webui.py @@ -272,8 +272,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", @@ -287,6 +288,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( @@ -341,6 +343,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) generated = save_dir / "generated.wav" @@ -385,7 +388,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, @@ -873,6 +876,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 index 8e7b696..45936f7 100644 --- a/webui_svc.py +++ b/webui_svc.py @@ -132,8 +132,9 @@ def _tips_md() -> 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", @@ -148,6 +149,7 @@ class AppState: 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]: @@ -203,6 +205,7 @@ class AppState: 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) @@ -244,7 +247,7 @@ class AppState: return False, f"svc inference failed: {e}", None -APP_STATE = AppState() +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): @@ -448,6 +451,7 @@ if __name__ == "__main__": 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()