add fp16 support for svs and svc inference
This commit is contained in:
+17
-2
@@ -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)
|
||||
|
||||
+30
-5
@@ -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)
|
||||
|
||||
+2
-1
@@ -25,4 +25,5 @@ python -m cli.inference \
|
||||
--phoneset_path $phoneset_path \
|
||||
--save_dir $save_dir \
|
||||
--auto_shift \
|
||||
--pitch_shift 0
|
||||
--pitch_shift 0 \
|
||||
--fp16
|
||||
@@ -24,4 +24,5 @@ python -m cli.inference_svc \
|
||||
--target_f0_path $target_f0_path \
|
||||
--save_dir $save_dir \
|
||||
--auto_shift \
|
||||
--pitch_shift 0
|
||||
--pitch_shift 0 \
|
||||
--fp16
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
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)
|
||||
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, :]
|
||||
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_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()
|
||||
|
||||
return generated_audio
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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()
|
||||
|
||||
+6
-2
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user