optimize GPU memory usage in vocal separation
This commit is contained in:
@@ -54,7 +54,7 @@ def build_model(args):
|
|||||||
return model, config
|
return model, config
|
||||||
|
|
||||||
|
|
||||||
def build_models(dict_args):
|
def build_models(dict_args, use_der: bool = False):
|
||||||
args = parse_args_inference(dict_args)
|
args = parse_args_inference(dict_args)
|
||||||
|
|
||||||
########## load model ##########
|
########## load model ##########
|
||||||
@@ -65,25 +65,26 @@ def build_models(dict_args):
|
|||||||
|
|
||||||
sep_model, sep_config = build_model(args)
|
sep_model, sep_config = build_model(args)
|
||||||
|
|
||||||
args.config_path = args.der_config_path
|
if use_der:
|
||||||
args.start_check_point = args.der_start_check_point
|
args.config_path = args.der_config_path
|
||||||
|
args.start_check_point = args.der_start_check_point
|
||||||
dereverb_model, dereverb_config = build_model(args)
|
dereverb_model, dereverb_config = build_model(args)
|
||||||
|
else:
|
||||||
sep_model = sep_model
|
dereverb_model, dereverb_config = None, None
|
||||||
dereverb_model = dereverb_model
|
|
||||||
|
|
||||||
return sep_model, sep_config, dereverb_model, dereverb_config, args
|
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):
|
def main(args, sep_model=None, sep_config=None, dereverb_model=None, dereverb_config=None, device=None):
|
||||||
|
|
||||||
######## process data ##########
|
######## process data ##########
|
||||||
sample_rate = getattr(sep_config.audio, 'sample_rate', 44100)
|
sample_rate = getattr(sep_config.audio, 'sample_rate', 44100)
|
||||||
path = args.input_path
|
path = args.input_path
|
||||||
|
|
||||||
mix, _ = librosa.load(path, sr=sample_rate, mono=False)
|
mix, _ = librosa.load(path, sr=sample_rate, mono=False)
|
||||||
vocals = process(mix, sep_model, args, sep_config, device)
|
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
|
accompaniment = mix - dereverbed_vocals
|
||||||
|
|
||||||
return mix, vocals, dereverbed_vocals, accompaniment, sample_rate
|
return mix, vocals, dereverbed_vocals, accompaniment, sample_rate
|
||||||
@@ -113,6 +114,8 @@ class VocalSeparator:
|
|||||||
der_model_path: str,
|
der_model_path: str,
|
||||||
der_config_path: str,
|
der_config_path: str,
|
||||||
*,
|
*,
|
||||||
|
chunk_length_sec: int = 5,
|
||||||
|
use_der: bool = False,
|
||||||
model_type: str = "mel_band_roformer",
|
model_type: str = "mel_band_roformer",
|
||||||
disable_detailed_pbar: bool = True,
|
disable_detailed_pbar: bool = True,
|
||||||
device: str = "cuda",
|
device: str = "cuda",
|
||||||
@@ -122,11 +125,14 @@ class VocalSeparator:
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
device: Torch device string, e.g. ``"cuda:0"``.
|
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.
|
model_type: Separation model type key.
|
||||||
sep_config_path: Config path for separation model.
|
sep_config_path: Config path for separation model.
|
||||||
sep_start_check_point: Checkpoint path for separation model.
|
sep_start_check_point: Checkpoint path for separation model.
|
||||||
der_config_path: Config path for dereverb model.
|
der_config_path: Config path for dereverb model.
|
||||||
der_start_check_point: Checkpoint 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.
|
disable_detailed_pbar: Disable detailed progress bars in underlying utils.
|
||||||
verbose: Whether to print verbose logs.
|
verbose: Whether to print verbose logs.
|
||||||
"""
|
"""
|
||||||
@@ -144,10 +150,15 @@ class VocalSeparator:
|
|||||||
if verbose:
|
if verbose:
|
||||||
print("[vocal extraction] init: start")
|
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)
|
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_model = sep_model
|
||||||
self.sep_config = sep_config
|
self.sep_config = sep_config
|
||||||
@@ -158,8 +169,9 @@ class VocalSeparator:
|
|||||||
self.verbose = verbose
|
self.verbose = verbose
|
||||||
|
|
||||||
if verbose:
|
if verbose:
|
||||||
|
der_status = "loaded" if dereverb_model is not None else "skipped"
|
||||||
print(
|
print(
|
||||||
"[vocal extraction] init success: sep=loaded, dereverb=loaded, device=",
|
"[vocal extraction] init success: sep=loaded, dereverb=%s, device=" % der_status,
|
||||||
device,
|
device,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user