mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
hubert path param
This commit is contained in:
@@ -117,11 +117,15 @@ parser = argparse.ArgumentParser(description='Run VC inference')
|
||||
|
||||
parser.add_argument('--model_path', type=str, help='path to the .pth file', required=True)
|
||||
parser.add_argument('--model_index_path', type=str, help='path to the .index file (empty string for none)', required=False, default='')
|
||||
parser.add_argument('--hubert_model_path', type=str, help='path to the hubert model .pth file (default hubert_base.pt)',
|
||||
required=False, default='hubert_base.pt')
|
||||
|
||||
parser.add_argument('--input_audio_filename', type=str, help='source audio file', required=True)
|
||||
parser.add_argument('--output_audio_filename', type=str, help='result audio file', required=True)
|
||||
#parser.add_argument('--output_metadata_filename', type=str, help='where to save extra metadata', required=True)
|
||||
|
||||
########
|
||||
# TODO: determine what we should do with these
|
||||
########
|
||||
parser.add_argument('--f0_method', type=str, help='harvest (default) or pm', required=False, default='harvest')
|
||||
parser.add_argument('--f0_up_key', type=str, help='???? (default "0" ???)', required=False, default='0')
|
||||
parser.add_argument('--index_rate', type=float, help='???? (default 0.66 ???)', required=False, default=0.66)
|
||||
@@ -130,6 +134,7 @@ parser.add_argument('--resample_sr', type=int, help='???? (default 0 ???)', requ
|
||||
parser.add_argument('--rms_mix_rate', type=float, help='???? (default 1 ???)', required=False, default=1.0)
|
||||
parser.add_argument('--protect', type=float, help='???? (default 0.33 ???)', required=False, default=0.33)
|
||||
|
||||
# Device / precision
|
||||
parser.add_argument('--device', type=str, help='compute device (eg. "cuda:0"), default "cuda:0".', required=False, default='cuda:0')
|
||||
parser.add_argument('--is_half', type=bool, help='halfwidth tensors (default false)', required=False, default=False)
|
||||
|
||||
@@ -148,7 +153,7 @@ filter_radius = args.filter_radius
|
||||
resample_sr = args.resample_sr
|
||||
rms_mix_rate = args.rms_mix_rate
|
||||
protect = args.protect
|
||||
|
||||
hubert_path = args.hubert_model_path
|
||||
|
||||
config = Config(device, is_half)
|
||||
now_dir = os.getcwd()
|
||||
@@ -167,10 +172,10 @@ from scipy.io import wavfile
|
||||
hubert_model = None
|
||||
|
||||
|
||||
def load_hubert():
|
||||
def load_hubert(hubert_path="hubert_base.pt"):
|
||||
global hubert_model
|
||||
models, saved_cfg, task = checkpoint_utils.load_model_ensemble_and_task(
|
||||
["hubert_base.pt"],
|
||||
[hubert_path],
|
||||
suffix="",
|
||||
)
|
||||
hubert_model = models[0]
|
||||
@@ -182,7 +187,7 @@ def load_hubert():
|
||||
hubert_model.eval()
|
||||
|
||||
|
||||
def vc_single(sid, input_audio, f0_up_key, f0_file, f0_method, file_index, index_rate):
|
||||
def vc_single(sid, input_audio, f0_up_key, f0_file, f0_method, file_index, index_rate, hubert_path="hubert_base.pt"):
|
||||
global tgt_sr, net_g, vc, hubert_model, version
|
||||
if input_audio is None:
|
||||
return "You need to upload an audio", None
|
||||
@@ -190,7 +195,7 @@ def vc_single(sid, input_audio, f0_up_key, f0_file, f0_method, file_index, index
|
||||
audio = load_audio(input_audio, 16000)
|
||||
times = [0, 0, 0]
|
||||
if hubert_model == None:
|
||||
load_hubert()
|
||||
load_hubert(hubert_path)
|
||||
if_f0 = cpt.get("f0", 1)
|
||||
# audio_opt=vc.pipeline(hubert_model,net_g,sid,audio,times,f0_up_key,f0_method,file_index,file_big_npy,index_rate,if_f0,f0_file=f0_file)
|
||||
audio_opt = vc.pipeline(
|
||||
@@ -251,7 +256,7 @@ get_vc(model_path)
|
||||
|
||||
file_path = input_path
|
||||
wav_opt = vc_single(
|
||||
0, file_path, f0up_key, None, f0method, index_path, index_rate
|
||||
0, file_path, f0up_key, None, f0method, index_path, index_rate, hubert_path
|
||||
)
|
||||
out_path = opt_path
|
||||
wavfile.write(out_path, tgt_sr, wav_opt)
|
||||
|
||||
Reference in New Issue
Block a user