From acd498ad9e36415a6d27126216418bddb3414cb6 Mon Sep 17 00:00:00 2001 From: jlqian98 Date: Mon, 16 Mar 2026 13:21:08 +0800 Subject: [PATCH] optimize GPU memory usage in vocal separation --- preprocess/tools/vocal_separation/model.py | 38 ++++++++++++++-------- 1 file changed, 25 insertions(+), 13 deletions(-) diff --git a/preprocess/tools/vocal_separation/model.py b/preprocess/tools/vocal_separation/model.py index ce6679c..9fcefe1 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, )