added rvc model training
This commit is contained in:
@@ -0,0 +1,4 @@
|
||||
from pathlib import Path
|
||||
import os
|
||||
|
||||
BASE_DIR = Path(os.path.dirname(os.path.abspath(__file__))).parent
|
||||
@@ -1,9 +1,7 @@
|
||||
|
||||
from io import BytesIO
|
||||
import os
|
||||
import shutil
|
||||
import numpy as np
|
||||
from pytube import YouTube
|
||||
import yt_dlp
|
||||
import torch
|
||||
from .settings import MERGE_OPTIONS
|
||||
@@ -33,7 +31,8 @@ def to_audio_dict(audio, sr):
|
||||
class LoadAudio:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
input_dir = input_path
|
||||
input_dir = os.path.join(input_path,"audio")
|
||||
os.makedirs(input_dir,exist_ok=True)
|
||||
files = get_filenames(root=input_dir,exts=SUPPORTED_AUDIO,format_func=os.path.basename)
|
||||
|
||||
return {
|
||||
@@ -49,7 +48,7 @@ class LoadAudio:
|
||||
FUNCTION = "load_audio"
|
||||
|
||||
def load_audio(self, audio, sr):
|
||||
audio_path = os.path.join(input_path,audio) #folder_paths.get_annotated_filepath(audio)
|
||||
audio_path = os.path.join(input_path,"audio",audio)
|
||||
widgetId = get_hash(audio_path)
|
||||
audio_name = os.path.basename(audio).split(".")[0]
|
||||
sr = None if sr=="None" else int(sr)
|
||||
@@ -90,7 +89,9 @@ class DownloadAudio:
|
||||
widgetId = get_hash(url, sr)
|
||||
sr = None if sr=="None" else int(sr)
|
||||
audio_name = widgetId if song_name=="" else song_name
|
||||
audio_path = os.path.join(input_path,f"{audio_name}.{format}")
|
||||
input_dir = os.path.join(input_path,"audio")
|
||||
os.makedirs(input_dir,exist_ok=True)
|
||||
audio_path = os.path.join(input_dir,f"{audio_name}.{format}")
|
||||
|
||||
if os.path.isfile(audio_path): input_audio = load_input_audio(audio_path,sr=sr)
|
||||
else:
|
||||
@@ -135,6 +136,7 @@ class MergeAudioNode:
|
||||
|
||||
RETURN_TYPES = ("VHS_AUDIO","AUDIO")
|
||||
RETURN_NAMES = ("vhs_audio","audio")
|
||||
OUTPUT_NODE = True
|
||||
|
||||
FUNCTION = "merge"
|
||||
|
||||
@@ -143,7 +145,7 @@ class MergeAudioNode:
|
||||
def merge(self, audio1, audio2, sr="None", merge_type="median", normalize=False, audio3_opt=None, audio4_opt=None):
|
||||
|
||||
input_audios = [get_audio(audio) for audio in [audio1, audio2, audio3_opt, audio4_opt] if audio is not None]
|
||||
widgetId = get_hash(*input_audios,sr,merge_type,normalize)
|
||||
widgetId = get_hash(*[audio_to_bytes(*audio) for audio in input_audios],sr,merge_type,normalize)
|
||||
audio_path = os.path.join(temp_path,"preview",f"{widgetId}.flac")
|
||||
|
||||
if os.path.isfile(audio_path): merged_audio = load_input_audio(audio_path)
|
||||
@@ -271,7 +273,6 @@ class AudioBatchValueNode:
|
||||
if inverse: x_norm=1-x_norm
|
||||
x_norm = x_norm * output_range + output_min
|
||||
|
||||
# batch_value = ",\n".join([f'{n}:({v})' for n,v in enumerate(x_norm)])
|
||||
if print_output:
|
||||
print(f"{audio_rms.min()=} {audio_rms.max()=} {audio_rms.mean()=} {len(audio_rms)=}")
|
||||
print(f"{x_norm.min()=} {x_norm.max()=} {x_norm.mean()=} {len(x_norm)=}")
|
||||
@@ -297,7 +298,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"RVC-Studio.LoadAudio": "🌺Load Audio",
|
||||
"DownloadAudio": "🌺Youtube Downloader",
|
||||
"RVC-Studio.PreviewAudio": "🌺Preview Audio",
|
||||
"RVC-Studio.PreviewAudio": "🌺Save Audio",
|
||||
"MergeAudioNode": "🌺Merge Audio",
|
||||
"AudioBatchValueNode": "🌺Audio RMS Batch Values",
|
||||
}
|
||||
+326
-31
@@ -1,29 +1,45 @@
|
||||
import json
|
||||
import multiprocessing
|
||||
import os
|
||||
import shutil
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from . import BASE_DIR
|
||||
|
||||
from ..lib.train.utils import HParams
|
||||
|
||||
from ..training_cli import train_model
|
||||
|
||||
from ..preprocessing_utils import extract_features_trainset, preprocess_trainset
|
||||
|
||||
from .audio_nodes import get_audio
|
||||
from ..config import config
|
||||
|
||||
from ..lib.model_utils import load_hubert
|
||||
|
||||
from .utils import model_downloader
|
||||
from .settings import PITCH_EXTRACTION_OPTIONS
|
||||
from .settings.downloader import RVC_DOWNLOAD_LINK, RVC_INDEX, RVC_MODELS, download_file
|
||||
from .utils import MultipleTypeProxy, increment_filename_no_overwrite, model_downloader
|
||||
from .settings import PITCH_EXTRACTION_OPTIONS, SR_MAP
|
||||
from .settings.downloader import PRETRAINED_MODELS, RVC_DOWNLOAD_LINK, RVC_INDEX, RVC_MODELS, download_file, extract_zip_without_structure
|
||||
|
||||
from ..lib.audio import SUPPORTED_AUDIO, audio_to_bytes, bytes_to_audio, load_input_audio, save_input_audio
|
||||
from ..lib.audio import SUPPORTED_AUDIO, audio_to_bytes, load_input_audio, save_input_audio
|
||||
|
||||
from ..vc_infer_pipeline import get_vc, vc_single
|
||||
import folder_paths
|
||||
from ..lib.utils import get_filenames, get_hash, get_optimal_torch_device
|
||||
from ..lib.utils import get_filenames, get_hash, get_optimal_threads, get_optimal_torch_device
|
||||
from ..lib import BASE_CACHE_DIR, BASE_MODELS_DIR
|
||||
|
||||
input_path = folder_paths.get_input_directory()
|
||||
temp_path = folder_paths.get_temp_directory()
|
||||
cache_dir = os.path.join(BASE_CACHE_DIR,"rvc")
|
||||
# output_path = folder_paths.get_output_directory()
|
||||
# base_path = os.path.dirname(input_path)
|
||||
output_path = folder_paths.get_output_directory()
|
||||
node_path = os.path.join(BASE_MODELS_DIR,"custom_nodes/ComfyUI-UVR5")
|
||||
weights_path = os.path.join(BASE_MODELS_DIR, "uvr5")
|
||||
dataset_path = os.path.join(input_path,"datasets")
|
||||
device = get_optimal_torch_device()
|
||||
CATEGORY = "🌺RVC-Studio/rvc"
|
||||
LOG_DIR = os.path.join(BASE_DIR,"logs")
|
||||
|
||||
class LoadPitchExtractionParams:
|
||||
@classmethod
|
||||
@@ -37,7 +53,6 @@ class LoadPitchExtractionParams:
|
||||
"min": 0., #Minimum value
|
||||
"max": 1., #Maximum value
|
||||
"step": .01, #Slider's step
|
||||
"display": "slider"
|
||||
}),
|
||||
"resample_sr": ([0,16000,32000,40000,44100,48000],{"default": 0}),
|
||||
"rms_mix_rate": ("FLOAT",{
|
||||
@@ -45,14 +60,12 @@ class LoadPitchExtractionParams:
|
||||
"min": 0., #Minimum value
|
||||
"max": 1., #Maximum value
|
||||
"step": .01, #Slider's step
|
||||
"display": "slider"
|
||||
}),
|
||||
"protect": ("FLOAT",{
|
||||
"default": 0.25,
|
||||
"min": 0., #Minimum value
|
||||
"max": .5, #Maximum value
|
||||
"step": .01, #Slider's step
|
||||
"display": "slider"
|
||||
})
|
||||
},
|
||||
}
|
||||
@@ -67,10 +80,6 @@ class LoadPitchExtractionParams:
|
||||
def load_params(self, **params):
|
||||
if "rmvpe" in params.get("f0_method",""): model_downloader("rmvpe.pt")
|
||||
return (params,)
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, **params):
|
||||
return get_hash(**params)
|
||||
|
||||
class LoadHubertModel:
|
||||
@classmethod
|
||||
@@ -96,17 +105,13 @@ class LoadHubertModel:
|
||||
hubert_model = lambda:load_hubert(model_path,config=config)
|
||||
return (hubert_model,)
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, model):
|
||||
return get_hash(model)
|
||||
|
||||
class LoadRVCModelNode:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
model_list = RVC_MODELS + get_filenames(root=BASE_MODELS_DIR,folder="RVC",exts=["pth"],format_func=lambda x: f"RVC/{os.path.basename(x)}")
|
||||
model_list = list(set(model_list)) # dedupe
|
||||
index_list = ["None"] + RVC_INDEX + get_filenames(root=BASE_MODELS_DIR,folder="RVC",exts=["index"],format_func=lambda x: f"RVC/.index/{os.path.basename(x)}")
|
||||
index_list = [""] + RVC_INDEX + get_filenames(root=os.path.join(BASE_MODELS_DIR,"RVC"),folder=".index",exts=["index"],format_func=lambda x: f"RVC/.index/{os.path.basename(x)}")
|
||||
index_list = list(set(index_list)) # dedupe
|
||||
|
||||
return {
|
||||
@@ -114,7 +119,7 @@ class LoadRVCModelNode:
|
||||
'model': (model_list,{"default": model_list[0]}),
|
||||
},
|
||||
"optional": {
|
||||
"index": (index_list,{"default": "None"}),
|
||||
"index": (index_list,{"default": ""}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -125,8 +130,7 @@ class LoadRVCModelNode:
|
||||
|
||||
FUNCTION = 'load_model'
|
||||
|
||||
|
||||
def load_model(self, model, index="None"):
|
||||
def load_model(self, model, index=""):
|
||||
model_path = file_index = None
|
||||
try:
|
||||
filename = os.path.basename(model)
|
||||
@@ -137,7 +141,7 @@ class LoadRVCModelNode:
|
||||
download_link = f"{RVC_DOWNLOAD_LINK}{model}"
|
||||
if download_file((model_path, download_link)): print(f"successfully downloaded: {model_path}")
|
||||
|
||||
if not index=="None":
|
||||
if index:
|
||||
file_index = os.path.join(BASE_MODELS_DIR,subfolder,".index",os.path.basename(index))
|
||||
if not os.path.isfile(file_index):
|
||||
download_link = f"{RVC_DOWNLOAD_LINK}{index}"
|
||||
@@ -146,10 +150,6 @@ class LoadRVCModelNode:
|
||||
print(f"Error in {self.__class__.__name__}: {e}")
|
||||
raise e
|
||||
finally: return (lambda:get_vc(model_path, file_index),filename.split(".")[0])
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, model, index):
|
||||
return get_hash(model, index)
|
||||
|
||||
class RVCNode:
|
||||
|
||||
@@ -162,7 +162,7 @@ class RVCNode:
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"audio": ("VHS_AUDIO",),
|
||||
"audio": (MultipleTypeProxy('AUDIO,VHS_AUDIO'),),
|
||||
"model": ("RVC_MODEL",),
|
||||
"hubert_model": ("HUBERT_MODEL",),
|
||||
"pitch_extraction_params": ("PITCH_EXTRACTION",),
|
||||
@@ -190,13 +190,15 @@ class RVCNode:
|
||||
|
||||
def convert(self, audio, model, hubert_model, pitch_extraction_params, f0_up_key, format="flac", use_cache=True):
|
||||
|
||||
widgetId = get_hash(audio(), model, f0_up_key, *pitch_extraction_params.items())
|
||||
input_audio = get_audio(audio)
|
||||
voice_model = model()
|
||||
feature_model = hubert_model()
|
||||
widgetId = get_hash(feature_model, f0_up_key, audio_to_bytes(*input_audio), *voice_model.items(), *pitch_extraction_params.items())
|
||||
cache_name = os.path.join(BASE_CACHE_DIR,"rvc",f"{widgetId}.{format}")
|
||||
|
||||
if use_cache and os.path.isfile(cache_name): output_audio = load_input_audio(cache_name)
|
||||
else:
|
||||
input_audio = bytes_to_audio(audio())
|
||||
output_audio = vc_single(hubert_model=hubert_model(),input_audio=input_audio,f0_up_key=f0_up_key,**model(),**pitch_extraction_params)
|
||||
output_audio = vc_single(hubert_model=feature_model,input_audio=input_audio,f0_up_key=f0_up_key,**voice_model,**pitch_extraction_params)
|
||||
|
||||
if use_cache:
|
||||
print(save_input_audio(cache_name, output_audio))
|
||||
@@ -209,6 +211,295 @@ class RVCNode:
|
||||
if not os.path.isfile(preview_file): shutil.copyfile(cache_name,preview_file)
|
||||
return {"ui": {"preview": [{"filename": audio_name, "type": "temp", "subfolder": "preview", "widgetId": widgetId}]}, "result": (lambda:audio_to_bytes(*output_audio),)}
|
||||
|
||||
class RVCProcessDatasetNode:
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
|
||||
os.makedirs(dataset_path, exist_ok=True)
|
||||
DATASETS = [""] + [dname for dname in os.listdir(dataset_path) if dname.endswith("zip")]
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"model_name": ("STRING",{"default": ""}),
|
||||
"dataset": (DATASETS, {"default": ""}),
|
||||
"hubert_model": ("HUBERT_MODEL",),
|
||||
},
|
||||
"optional": {
|
||||
"pitch_extraction_params": ("PITCH_EXTRACTION",{"default": {}}),
|
||||
"sr": (["32k","40k","48k"], {"default": "40k"}),
|
||||
"n_threads": ("INT", {"default": get_optimal_threads(), "min": 1, "max": multiprocessing.cpu_count()}),
|
||||
"period": ("FLOAT", {"default": 3., "min": 1., "max": 10., "step": .1}),
|
||||
"overlap": ("FLOAT",{"default": .3, "min": .1, "max": 1., "step": .1}),
|
||||
"max_volume": ("FLOAT",{"default": 1., "min": .1, "max": 1., "step": .05}),
|
||||
"alpha": ("FLOAT",{"default": .75, "min": .05, "max": 1., "step": .05}),
|
||||
"mute_ratio": ("FLOAT",{"default": .0, "min": .0, "max": .5, "step": .01}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("RVC_DATASET_PIPE",)
|
||||
RETURN_NAMES = ("rvc_dataset_pipe",)
|
||||
|
||||
FUNCTION = "process"
|
||||
|
||||
CATEGORY = CATEGORY
|
||||
|
||||
def process(self, model_name: str, dataset: str, hubert_model, pitch_extraction_params={}, sr="40k", n_threads=1, period=3., overlap=.3, max_volume=1., alpha=.75, mute_ratio=.0):
|
||||
|
||||
assert model_name, "Please provide a model name!"
|
||||
assert dataset, "Please upload a dataset!"
|
||||
|
||||
f0_method = pitch_extraction_params.get("f0_method")
|
||||
cache_name = get_hash(model_name, dataset, period, overlap, max_volume, alpha, mute_ratio, f0_method)
|
||||
model_log_dir = os.path.join(output_path,"logs",cache_name)
|
||||
os.makedirs(model_log_dir,exist_ok=True)
|
||||
|
||||
filelist_path = os.path.join(model_log_dir, "filelist.txt")
|
||||
|
||||
if not os.path.isfile(filelist_path):
|
||||
dataset_dir = os.path.join(input_path,"datasets",cache_name)
|
||||
|
||||
if dataset.endswith("zip"):
|
||||
files = extract_zip_without_structure(os.path.join(dataset_path,dataset),dataset_dir)
|
||||
assert len(files), "Failed to extract zip file..."
|
||||
|
||||
print(preprocess_trainset(dataset_dir,SR_MAP[sr],n_threads,model_log_dir,period,overlap,max_volume,alpha))
|
||||
|
||||
print(extract_features_trainset(hubert_model(), model_log_dir,n_p=n_threads,f0method=f0_method,device=device,if_f0=bool(f0_method),version="v2"))
|
||||
|
||||
gt_wavs_dir = os.path.join(model_log_dir,"0_gt_wavs")
|
||||
feature_dir = os.path.join(model_log_dir,"3_feature768")
|
||||
os.makedirs(gt_wavs_dir, exist_ok=True)
|
||||
os.makedirs(feature_dir, exist_ok=True)
|
||||
|
||||
# add training data
|
||||
if f0_method:
|
||||
f0_dir = os.path.join(model_log_dir,"2a_f0")
|
||||
f0nsf_dir = os.path.join(model_log_dir,"2b-f0nsf")
|
||||
names = (
|
||||
set([os.path.splitext(name)[0] for name in os.listdir(feature_dir)])
|
||||
& set([os.path.splitext(name)[0] for name in os.listdir(f0_dir)])
|
||||
& set([os.path.splitext(name)[0] for name in os.listdir(f0nsf_dir)])
|
||||
)
|
||||
else:
|
||||
names = set(
|
||||
[os.path.splitext(name)[0] for name in os.listdir(feature_dir)]
|
||||
)
|
||||
opt = []
|
||||
missing_data = []
|
||||
for name in names:
|
||||
name_parts = name.split(",")
|
||||
gt_name = name if len(name_parts) == 1 else name_parts[-1]
|
||||
gt_file = os.path.join(gt_wavs_dir,gt_name)
|
||||
if not os.path.isfile(gt_file):
|
||||
print(f"{gt_name} not found!")
|
||||
missing_data.append(gt_name)
|
||||
continue #skip data
|
||||
|
||||
if f0_method:
|
||||
data = "|".join([
|
||||
gt_file,
|
||||
os.path.join(feature_dir,f"{name}.npy"),
|
||||
os.path.join(f0_dir,f"{name}.npy"),
|
||||
os.path.join(f0nsf_dir,f"{name}.npy"),
|
||||
str(0)
|
||||
])
|
||||
else:
|
||||
data = "|".join([
|
||||
gt_file,
|
||||
os.path.join(feature_dir,f"{name}.npy"),
|
||||
str(0)
|
||||
])
|
||||
opt.append(data)
|
||||
|
||||
assert len(missing_data)==0, f"missing ground truth data: {len(opt)=}, {len(missing_data)=}"
|
||||
|
||||
# add mute data
|
||||
fea_dim = 768
|
||||
num_mute = max(2,int(len(opt)*mute_ratio)) # use 1% mute file or 2 copies (like original repo)
|
||||
for _ in range(num_mute):
|
||||
if f0_method:
|
||||
data = "|".join([
|
||||
os.path.join(LOG_DIR,"mute","0_gt_wavs",f"mute{sr}.wav"),
|
||||
os.path.join(LOG_DIR,"mute",f"3_feature{fea_dim}","mute.npy"),
|
||||
os.path.join(LOG_DIR,"mute","2a_f0","mute.wav.npy"),
|
||||
os.path.join(LOG_DIR,"mute","2b-f0nsf","mute.wav.npy"),
|
||||
str(0)
|
||||
])
|
||||
else:
|
||||
data = "|".join([
|
||||
os.path.join(LOG_DIR,"mute","0_gt_wavs",f"mute{sr}.wav"),
|
||||
os.path.join(LOG_DIR,"mute",f"3_feature{fea_dim}","mute.npy"),
|
||||
str(0)
|
||||
])
|
||||
opt.append(data)
|
||||
|
||||
|
||||
np.random.shuffle(opt)
|
||||
with open(filelist_path, "w") as f:
|
||||
f.write("\n".join(opt))
|
||||
print("write filelist done")
|
||||
return (dict(
|
||||
sample_rate=sr,
|
||||
model_dir=model_log_dir,
|
||||
name=model_name,
|
||||
training_files=filelist_path,
|
||||
if_f0=bool(f0_method),
|
||||
pitch_extraction_params=pitch_extraction_params,
|
||||
hubert_model=hubert_model
|
||||
),)
|
||||
|
||||
|
||||
class RVCTrainModelNode:
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
DEVICES = [str(i) for i in range(torch.cuda.device_count())]
|
||||
PRETRAINED_G = [model for model in PRETRAINED_MODELS if "G" in model] + get_filenames(root=BASE_MODELS_DIR,folder="pretrained_v2",name_filters=["G"],format_func=lambda x: f"pretrained_v2/{os.path.basename(x)}")
|
||||
PRETRAINED_G = list(set(PRETRAINED_G)) # dedupe
|
||||
PRETRAINED_D = [model for model in PRETRAINED_MODELS if "D" in model] + get_filenames(root=BASE_MODELS_DIR,folder="pretrained_v2",name_filters=["D"],format_func=lambda x: f"pretrained_v2/{os.path.basename(x)}")
|
||||
PRETRAINED_D = list(set(PRETRAINED_D)) # dedupe
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"rvc_dataset_pipe": ("RVC_DATASET_PIPE",),
|
||||
},
|
||||
"optional": dict(
|
||||
gpu=(DEVICES, {"default": DEVICES[0]}),
|
||||
batch_size=("INT",dict(default=4,min=1,max=64)),
|
||||
total_epoch=("INT",dict(default=100,min=10,max=1000,step=10)),
|
||||
save_every_epoch=("INT",dict(default=0,min=0,max=100)),
|
||||
pretrained_G=(PRETRAINED_G,{"default": PRETRAINED_G[0]}),
|
||||
pretrained_D=(PRETRAINED_D,{"default": PRETRAINED_D[0]}),
|
||||
if_save_latest=("BOOLEAN",{"default":True}),
|
||||
if_cache_gpu=("BOOLEAN",{"default":True}),
|
||||
if_save_every_weights=("BOOLEAN",{"default":False}),
|
||||
train_index=("BOOLEAN",{"default": True}),
|
||||
rerun=("BOOLEAN",{"default": False}),
|
||||
)
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('RVC_MODEL', 'STRING', "HUBERT_MODEL", "PITCH_EXTRACTION" )
|
||||
RETURN_NAMES = ('model', 'model_name', "hubert_model", "pitch_extraction_params" )
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
FUNCTION = "train_model"
|
||||
|
||||
CATEGORY = CATEGORY
|
||||
|
||||
def train_model(self,
|
||||
rvc_dataset_pipe,
|
||||
gpu="0",
|
||||
batch_size=4,
|
||||
total_epoch=100,
|
||||
save_every_epoch=0,
|
||||
pretrained_G="",
|
||||
pretrained_D="",
|
||||
if_save_latest=True,
|
||||
if_cache_gpu=True,
|
||||
if_save_every_weights=False,
|
||||
train_index=True,
|
||||
rerun=False):
|
||||
|
||||
sample_rate = rvc_dataset_pipe["sample_rate"]
|
||||
name = rvc_dataset_pipe["name"]
|
||||
model_dir = rvc_dataset_pipe["model_dir"]
|
||||
if_f0 = rvc_dataset_pipe["if_f0"]
|
||||
|
||||
config_path = os.path.join(BASE_DIR,"configs",f"{sample_rate}{'' if sample_rate=='40k' else '_v2'}.json")
|
||||
with open(config_path,"r") as f:
|
||||
config = json.load(f)
|
||||
|
||||
hparams = HParams(**config)
|
||||
hparams.model_dir = hparams.experiment_dir = model_dir
|
||||
hparams.save_every_epoch = save_every_epoch
|
||||
hparams.name = name
|
||||
hparams.total_epoch = total_epoch
|
||||
hparams.pretrainG = model_downloader(pretrained_G)
|
||||
hparams.pretrainD = model_downloader(pretrained_D)
|
||||
hparams.version = "v2"
|
||||
hparams.gpus = gpu
|
||||
hparams.train.batch_size = batch_size
|
||||
hparams.sample_rate = sample_rate
|
||||
hparams.if_f0 = if_f0
|
||||
hparams.if_latest = if_save_latest
|
||||
hparams.save_every_weights = if_save_every_weights
|
||||
hparams.if_cache_data_in_gpu = if_cache_gpu
|
||||
hparams.data.training_files = rvc_dataset_pipe["training_files"]
|
||||
|
||||
file_index = self.train_index(model_dir, sample_rate, name) if train_index else None
|
||||
model_path = os.path.join(BASE_MODELS_DIR,"RVC",f"{name}.pth")
|
||||
if os.path.isfile(model_path) and rerun: model_path = increment_filename_no_overwrite(model_path)
|
||||
hparams.model_path = model_path
|
||||
|
||||
if not os.path.isfile(model_path): train_model(hparams)
|
||||
|
||||
return (lambda: get_vc(model_path, file_index), name, rvc_dataset_pipe["hubert_model"], rvc_dataset_pipe["pitch_extraction_params"])
|
||||
|
||||
def train_index(self, model_log_dir, sr, name):
|
||||
|
||||
key = get_hash(model_log_dir, sr, name)
|
||||
index_file = os.path.join(BASE_MODELS_DIR,"RVC",".index",f"{name}_v2_{sr}_{key}.index")
|
||||
|
||||
try:
|
||||
|
||||
if not os.path.isfile(index_file):
|
||||
from sklearn.cluster import MiniBatchKMeans
|
||||
import faiss
|
||||
|
||||
feature_dir = os.sep.join([model_log_dir, "3_feature768"])
|
||||
os.makedirs(feature_dir, exist_ok=True)
|
||||
|
||||
npys = []
|
||||
listdir_res = list(os.listdir(feature_dir))
|
||||
for fname in sorted(listdir_res):
|
||||
phone = np.load(os.path.join(feature_dir, fname))
|
||||
npys.append(phone)
|
||||
big_npy = np.concatenate(npys, 0)
|
||||
|
||||
big_npy_idx = np.arange(big_npy.shape[0])
|
||||
np.random.shuffle(big_npy_idx)
|
||||
big_npy = big_npy[big_npy_idx]
|
||||
|
||||
if big_npy.shape[0] > 2e5:
|
||||
big_npy = (
|
||||
MiniBatchKMeans(
|
||||
n_clusters=10000,
|
||||
verbose=True,
|
||||
batch_size=256 * config.n_cpu,
|
||||
compute_labels=False,
|
||||
init="random",
|
||||
)
|
||||
.fit(big_npy)
|
||||
.cluster_centers_
|
||||
)
|
||||
|
||||
n_ivf = min(int(16 * np.sqrt(big_npy.shape[0])), big_npy.shape[0] // 39)
|
||||
print(f"{big_npy.shape=} {n_ivf=}")
|
||||
index = faiss.index_factory(768, "IVF%s,Flat" % n_ivf)
|
||||
print("training index")
|
||||
index_ivf = faiss.extract_index_ivf(index) #
|
||||
index_ivf.nprobe = 1
|
||||
index.train(big_npy)
|
||||
print("adding index")
|
||||
batch_size_add = 8192
|
||||
for i in range(0, big_npy.shape[0], batch_size_add):
|
||||
index.add(big_npy[i : i + batch_size_add])
|
||||
|
||||
faiss.write_index(index,index_file)
|
||||
print(f"saved index file to {index_file}")
|
||||
return index_file
|
||||
except Exception as e:
|
||||
print(f"Failed to train index: {e}")
|
||||
return None
|
||||
|
||||
# A dictionary that contains all nodes you want to export with their names
|
||||
# NOTE: names should be globally unique
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
@@ -216,6 +507,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"RVCNode": RVCNode,
|
||||
"LoadHubertModel": LoadHubertModel,
|
||||
"LoadPitchExtractionParams": LoadPitchExtractionParams,
|
||||
"RVCProcessDatasetNode": RVCProcessDatasetNode,
|
||||
"RVCTrainModelNode": RVCTrainModelNode
|
||||
}
|
||||
|
||||
# A dictionary that contains the friendly/humanly readable titles for the nodes
|
||||
@@ -224,4 +517,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"RVCNode": "🌺Voice Changer",
|
||||
"LoadHubertModel": "🌺Load Hubert Model",
|
||||
"LoadPitchExtractionParams": "🌺Load Pitch Extraction Params",
|
||||
"RVCProcessDatasetNode": "🌺Process RVC Dataset",
|
||||
"RVCTrainModelNode": "🌺Train RVC Model",
|
||||
}
|
||||
@@ -1,11 +1,12 @@
|
||||
import platform
|
||||
import re
|
||||
import shutil
|
||||
import subprocess
|
||||
from typing import IO, List, Tuple
|
||||
import unicodedata
|
||||
import requests
|
||||
import os
|
||||
import zipfile
|
||||
from zipfile import ZipFile
|
||||
|
||||
from ...lib import BASE_CACHE_DIR, BASE_MODELS_DIR
|
||||
|
||||
@@ -31,7 +32,8 @@ RVC_MODELS = [
|
||||
RVC_INDEX = [
|
||||
"RVC/.index/added_IVF1063_Flat_nprobe_1_Sayano_v2.index",
|
||||
"RVC/.index/added_IVF985_Flat_nprobe_1_Fuji_v2.index",
|
||||
"RVC/.index/Monika_v2_40k.index"
|
||||
"RVC/.index/Monika_v2_40k.index",
|
||||
"RVC/.index/Sayano_v2_40k.index"
|
||||
]
|
||||
BASE_MODELS = ["content-vec-best.safetensors", "rmvpe.pt"]
|
||||
VITS_MODELS = ["VITS/pretrained_ljs.pth"]
|
||||
@@ -99,6 +101,28 @@ def save_file_generator(save_dir: str, data: List[IO]):
|
||||
data_path = os.path.join(save_dir,datum.name)
|
||||
yield (data_path, datum.read())
|
||||
|
||||
def extract_zip_without_structure(zip_path, extract_to, cleanup=False):
|
||||
os.makedirs(extract_to,exist_ok=True)
|
||||
try:
|
||||
with ZipFile(zip_path, 'r') as zip_ref:
|
||||
zipped_files = zip_ref.namelist()
|
||||
extracted_files = os.listdir(extract_to)
|
||||
if len(zipped_files)>len(extracted_files):
|
||||
for member in zipped_files:
|
||||
# Get the file name only (discard the directory structure)
|
||||
filename = os.path.basename(member)
|
||||
if filename: # Check if it's not an empty string
|
||||
# Create the full path for the extracted file
|
||||
file_path = os.path.join(extract_to, filename)
|
||||
# Extract the file
|
||||
with zip_ref.open(member) as source, open(file_path, 'wb') as target:
|
||||
shutil.copyfileobj(source, target)
|
||||
if cleanup: os.remove(zip_path) # cleanup
|
||||
print(f"Successfully extracted files to: {extract_to}")
|
||||
except Exception as error:
|
||||
print(f"Failed to extract files: {error}")
|
||||
finally: return os.listdir(extract_to)
|
||||
|
||||
def save_zipped_files(params: Tuple[str, any]):
|
||||
(data_path, datum) = params
|
||||
|
||||
@@ -113,14 +137,13 @@ def save_zipped_files(params: Tuple[str, any]):
|
||||
f.write(datum)
|
||||
|
||||
print(f"extracting zip file: {zip_path}")
|
||||
with zipfile.ZipFile(zip_path, 'r') as zip_ref:
|
||||
zip_ref.extractall(os.path.dirname(data_path))
|
||||
print(f"finished extracting zip file")
|
||||
files = extract_zip_without_structure(zip_path,os.path.dirname(data_path),cleanup=True)
|
||||
print(f"finished extracting {len(files)} files")
|
||||
|
||||
os.remove(zip_path) # cleanup
|
||||
return f"Successfully saved files to: {data_path}"
|
||||
return True
|
||||
except Exception as e:
|
||||
return f"Failed to save files: {e}"
|
||||
print(f"Failed to save files: {e}")
|
||||
return False
|
||||
|
||||
def slugify_filepath(filepath):
|
||||
# Split the path into directory and filename
|
||||
|
||||
@@ -6,7 +6,7 @@ from ..lib import BASE_MODELS_DIR
|
||||
|
||||
temp_path = folder_paths.get_temp_directory()
|
||||
|
||||
def model_downloader(model):
|
||||
def model_downloader(model: str) -> str:
|
||||
filename = os.path.basename(model)
|
||||
subfolder = os.path.dirname(model)
|
||||
model_path = os.path.join(BASE_MODELS_DIR,subfolder,filename)
|
||||
@@ -18,7 +18,7 @@ def model_downloader(model):
|
||||
|
||||
def increment_filename_no_overwrite(proposed_path):
|
||||
output_dir, filename_ext = os.path.split(proposed_path)
|
||||
filename, file_format = os.splitext(filename_ext)
|
||||
filename, file_format = os.path.splitext(filename_ext)
|
||||
files = os.listdir(output_dir)
|
||||
files = filter(lambda f: os.path.isfile(os.path.join(output_dir, f)), files)
|
||||
files = filter(lambda f: f.startswith(filename), files)
|
||||
@@ -26,7 +26,7 @@ def increment_filename_no_overwrite(proposed_path):
|
||||
file_numbers = [re.search(r'_(\d+)\.', f) for f in files]
|
||||
file_numbers = [int(f.group(1)) for f in file_numbers if f]
|
||||
this_file_number = max(file_numbers) + 1 if file_numbers else 1
|
||||
output_path = os.path.join(output_dir, filename + f'_{this_file_number}.' + file_format)
|
||||
output_path = os.path.join(output_dir, f'{filename}_{this_file_number}.{file_format}')
|
||||
return output_path
|
||||
|
||||
class MultipleTypeProxy(str):
|
||||
|
||||
+1
-1
@@ -65,7 +65,7 @@ class UVR5Node:
|
||||
if download_file(params): print(f"successfully downloaded: {model_path}")
|
||||
|
||||
input_audio = get_audio(audio)
|
||||
hash_name = get_hash(model, agg, format, *input_audio)
|
||||
hash_name = get_hash(model, agg, format, audio_to_bytes(*input_audio))
|
||||
audio_path = os.path.join(temp_path,"uvr",f"{hash_name}.wav")
|
||||
primary_path = os.path.join(cache_dir,hash_name,f"primary.{format}")
|
||||
secondary_path = os.path.join(cache_dir,hash_name,f"secondary.{format}")
|
||||
|
||||
@@ -278,22 +278,7 @@ def load_filepaths_and_text(filename, split="|"):
|
||||
|
||||
|
||||
def get_hparams(init=True):
|
||||
"""
|
||||
todo:
|
||||
结尾七人组:
|
||||
保存频率、总epoch done
|
||||
bs done
|
||||
pretrainG、pretrainD done
|
||||
卡号:os.en["CUDA_VISIBLE_DEVICES"] done
|
||||
if_latest done
|
||||
模型:if_f0 done
|
||||
采样率:自动选择config done
|
||||
是否缓存数据集进GPU:if_cache_data_in_gpu done
|
||||
|
||||
-m:
|
||||
自动决定training_files路径,改掉train_nsf_load_pretrain.py里的hps.data.training_files done
|
||||
-c不要了
|
||||
"""
|
||||
parser = argparse.ArgumentParser()
|
||||
# parser.add_argument('-c', '--config', type=str, default="configs/40k.json",help='JSON file for configuration')
|
||||
parser.add_argument(
|
||||
|
||||
@@ -41,7 +41,6 @@ class FeatureExtractor:
|
||||
"crepe-tiny": partial(self.get_f0_official_crepe_computation, model='model'),
|
||||
"mangio-crepe": self.get_f0_crepe_computation,
|
||||
"mangio-crepe-tiny": partial(self.get_f0_crepe_computation, model='model'),
|
||||
|
||||
}
|
||||
|
||||
def __del__(self):
|
||||
|
||||
+11
-19
@@ -1,6 +1,5 @@
|
||||
import sys, os, multiprocessing
|
||||
from threading import Thread
|
||||
from scipy import signal
|
||||
import numpy as np, os, traceback
|
||||
from .lib.model_utils import load_hubert
|
||||
from .lib.slicer2 import Slicer
|
||||
@@ -14,7 +13,7 @@ from .config import config
|
||||
import torch
|
||||
|
||||
class Preprocess:
|
||||
def __init__(self, sr, exp_dir, noparallel=True, period=3.0, overlap=.3, max_volume=.95):
|
||||
def __init__(self, sr, exp_dir, noparallel=True, period=3.0, overlap=.3, alpha=.75, max_volume=.95):
|
||||
self.slicer = Slicer(
|
||||
sr=sr,
|
||||
threshold=-42,
|
||||
@@ -24,12 +23,11 @@ class Preprocess:
|
||||
max_sil_kept=500
|
||||
)
|
||||
self.sr = sr
|
||||
self.bh, self.ah = signal.butter(N=5, Wn=48, btype="high", fs=self.sr)
|
||||
self.per = period
|
||||
self.overlap = overlap
|
||||
self.tail = self.per + self.overlap
|
||||
self.max = max_volume
|
||||
self.alpha = 0.75
|
||||
self.alpha = alpha
|
||||
self.exp_dir = exp_dir
|
||||
self.gt_wavs_dir = "%s/0_gt_wavs" % exp_dir
|
||||
self.wavs16k_dir = "%s/1_16k_wavs" % exp_dir
|
||||
@@ -70,10 +68,7 @@ class Preprocess:
|
||||
|
||||
def pipeline(self, path, idx0):
|
||||
try:
|
||||
audio = load_audio(path, self.sr)
|
||||
# zero phased digital filter cause pre-ringing noise...
|
||||
# audio = signal.filtfilt(self.bh, self.ah, audio)
|
||||
# audio = signal.lfilter(self.bh, self.ah, audio)
|
||||
audio,_ = load_input_audio(path, self.sr)
|
||||
|
||||
idx1 = 0
|
||||
for audio in self.slicer.slice(audio):
|
||||
@@ -121,7 +116,7 @@ class Preprocess:
|
||||
self.println("Fail. %s" % traceback.format_exc())
|
||||
|
||||
class FeatureInput(FeatureExtractor):
|
||||
def __init__(self, f0_method, exp_dir, samplerate=16000, hop_size=160, device="cpu", version="v2", if_f0=False):
|
||||
def __init__(self, model, f0_method, exp_dir, samplerate=16000, hop_size=160, device="cpu", version="v2", if_f0=False):
|
||||
self.sr = samplerate
|
||||
self.hop = hop_size
|
||||
self.f0_method = f0_method
|
||||
@@ -136,7 +131,7 @@ class FeatureInput(FeatureExtractor):
|
||||
self.f0_mel_min = 1127 * np.log(1 + self.f0_min / 700)
|
||||
self.f0_mel_max = 1127 * np.log(1 + self.f0_max / 700)
|
||||
|
||||
self.model = load_hubert(config)
|
||||
self.model = model
|
||||
|
||||
super().__init__(samplerate, config, onnx=False)
|
||||
|
||||
@@ -161,11 +156,8 @@ class FeatureInput(FeatureExtractor):
|
||||
"padding_mask": padding_mask.to(self.device),
|
||||
"output_layer": 9 if self.version == "v1" else 12, # layer 9
|
||||
}
|
||||
with torch.no_grad():
|
||||
logits = self.model.extract_features(**inputs)
|
||||
feats = (
|
||||
self.model.final_proj(logits[0]) if self.version == "v1" else logits[0]
|
||||
)
|
||||
|
||||
feats = self.model.extract_features(version=self.version,**inputs)
|
||||
|
||||
feats = feats.squeeze(0).float().cpu().numpy()
|
||||
if np.isnan(feats).sum() == 0:
|
||||
@@ -216,9 +208,9 @@ class FeatureInput(FeatureExtractor):
|
||||
except:
|
||||
self.printt("f0fail-%s-%s-%s" % (idx, inp_path, traceback.format_exc()))
|
||||
|
||||
def preprocess_trainset(inp_root, sr, n_p, exp_dir, period=3.0, overlap=.3):
|
||||
def preprocess_trainset(inp_root, sr, n_p, exp_dir, period=3.0, overlap=.3, max_volume=1., alpha=.75):
|
||||
try:
|
||||
pp = Preprocess(sr, exp_dir, period=period, overlap=overlap)
|
||||
pp = Preprocess(sr, exp_dir, period=period, overlap=overlap, max_volume=max_volume, alpha=alpha)
|
||||
pp.println("start preprocess")
|
||||
pp.println(sys.argv)
|
||||
pp.pipeline_mp_inp_dir(inp_root, n_p)
|
||||
@@ -229,9 +221,9 @@ def preprocess_trainset(inp_root, sr, n_p, exp_dir, period=3.0, overlap=.3):
|
||||
except Exception as e:
|
||||
return f"Failed to preprocess data: {e}"
|
||||
|
||||
def extract_features_trainset(exp_dir,n_p,f0method,device,version,if_f0):
|
||||
def extract_features_trainset(hubert_model,exp_dir,n_p,f0method,device,version,if_f0):
|
||||
try:
|
||||
featureInput = FeatureInput(f0_method=f0method,exp_dir=exp_dir,device=device,version=version,if_f0=if_f0)
|
||||
featureInput = FeatureInput(f0_method=f0method,exp_dir=exp_dir,device=device,version=version,if_f0=if_f0,model=hubert_model)
|
||||
paths = []
|
||||
inp_root = os.path.join(exp_dir,"1_16k_wavs")
|
||||
opt_root1 = os.path.join(exp_dir,"2a_f0")
|
||||
|
||||
+38
-38
@@ -7,9 +7,6 @@ from .lib import BASE_MODELS_DIR
|
||||
from .lib.train import utils
|
||||
import datetime
|
||||
|
||||
hps = utils.get_hparams()
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = hps.gpus.replace("-", ",")
|
||||
os.environ["NCCL_P2P_DISABLE"] = 1
|
||||
from random import shuffle, randint
|
||||
|
||||
import torch
|
||||
@@ -34,25 +31,13 @@ from .lib.train.data_utils import (
|
||||
DistributedBucketSampler,
|
||||
)
|
||||
|
||||
if hps.version == "v1":
|
||||
from .lib.infer_pack.models import (
|
||||
SynthesizerTrnMs256NSFsid as RVC_Model_f0,
|
||||
SynthesizerTrnMs256NSFsid_nono as RVC_Model_nof0,
|
||||
MultiPeriodDiscriminator,
|
||||
)
|
||||
else:
|
||||
from .lib.infer_pack.models import (
|
||||
SynthesizerTrnMs768NSFsid as RVC_Model_f0,
|
||||
SynthesizerTrnMs768NSFsid_nono as RVC_Model_nof0,
|
||||
MultiPeriodDiscriminatorV2 as MultiPeriodDiscriminator,
|
||||
)
|
||||
from .lib.train.losses import generator_loss, discriminator_loss, feature_loss, kl_loss
|
||||
from .lib.train.mel_processing import mel_spectrogram_torch, spec_to_mel_torch
|
||||
|
||||
global_step = 0
|
||||
least_loss = 40
|
||||
|
||||
def save_checkpoint(ckpt, sr, if_f0, name, epoch, version, hps, model_path=os.path.join(BASE_MODELS_DIR,"RVC")):
|
||||
def save_checkpoint(ckpt, name, epoch, hps, model_path=None):
|
||||
try:
|
||||
opt = OrderedDict()
|
||||
opt["weight"] = {}
|
||||
@@ -81,10 +66,11 @@ def save_checkpoint(ckpt, sr, if_f0, name, epoch, version, hps, model_path=os.pa
|
||||
hps.data.sampling_rate,
|
||||
]
|
||||
opt["info"] = "%sepoch" % epoch
|
||||
opt["sr"] = sr
|
||||
opt["f0"] = if_f0
|
||||
opt["version"] = version
|
||||
torch.save(opt, os.path.join(model_path,name+".pth"))
|
||||
opt["sr"] = hps.sample_rate
|
||||
opt["f0"] = hps.if_f0
|
||||
opt["version"] = hps.version
|
||||
if model_path is None: model_path=os.path.join(hps.model_dir,name+".pth")
|
||||
torch.save(opt, model_path)
|
||||
return "Success."
|
||||
except:
|
||||
return traceback.format_exc()
|
||||
@@ -101,8 +87,13 @@ class EpochRecorder:
|
||||
current_time = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
return f"[{current_time}] | ({elapsed_time_str})"
|
||||
|
||||
def train_model(hps: "utils.HParams"):
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = hps.gpus.replace("-", ",")
|
||||
os.environ["NCCL_P2P_DISABLE"] = "1"
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = str(randint(20000, 55555))
|
||||
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "max_split_size_mb:128,garbage_collection_threshold:0.8"
|
||||
|
||||
def main():
|
||||
n_gpus = len(hps.gpus.split("-")) if hps.gpus else torch.cuda.device_count()
|
||||
|
||||
if not torch.cuda.is_available() and torch.backends.mps.is_available():
|
||||
@@ -111,12 +102,12 @@ def main():
|
||||
# patch to unblock people without gpus. there is probably a better way.
|
||||
print("NO GPU DETECTED: falling back to CPU - this may take a while")
|
||||
n_gpus = 1
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = str(randint(20000, 55555))
|
||||
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "max_split_size_mb:128,garbage_collection_threshold:0.8"
|
||||
children = {}
|
||||
|
||||
gpu_devices = hps.gpus.split("-") if hps.gpus else range(n_gpus)
|
||||
|
||||
children = {}
|
||||
for i, device in enumerate(gpu_devices):
|
||||
mp.Pool
|
||||
subproc = mp.Process(
|
||||
target=run,
|
||||
args=(
|
||||
@@ -129,11 +120,23 @@ def main():
|
||||
children[i]=subproc
|
||||
subproc.start()
|
||||
|
||||
for i in gpu_devices:
|
||||
for i in children:
|
||||
children[i].join()
|
||||
|
||||
|
||||
def run(rank, n_gpus, hps, device):
|
||||
if hps.version == "v1":
|
||||
from .lib.infer_pack.models import (
|
||||
SynthesizerTrnMs256NSFsid as RVC_Model_f0,
|
||||
SynthesizerTrnMs256NSFsid_nono as RVC_Model_nof0,
|
||||
MultiPeriodDiscriminator,
|
||||
)
|
||||
else:
|
||||
from .lib.infer_pack.models import (
|
||||
SynthesizerTrnMs768NSFsid as RVC_Model_f0,
|
||||
SynthesizerTrnMs768NSFsid_nono as RVC_Model_nof0,
|
||||
MultiPeriodDiscriminatorV2 as MultiPeriodDiscriminator,
|
||||
)
|
||||
global global_step, least_loss
|
||||
if rank == 0:
|
||||
logger = utils.get_logger(hps.model_dir)
|
||||
@@ -538,13 +541,9 @@ def train_and_evaluate(
|
||||
loss_gen_all,
|
||||
save_checkpoint(
|
||||
ckpt,
|
||||
hps.sample_rate,
|
||||
hps.if_f0,
|
||||
f"loss{loss_gen_all:2.0f}_{hps.name}_e{epoch}",
|
||||
epoch,
|
||||
hps.version,
|
||||
hps,
|
||||
model_path=hps.model_dir
|
||||
hps
|
||||
),
|
||||
)
|
||||
)
|
||||
@@ -639,13 +638,9 @@ def train_and_evaluate(
|
||||
epoch,
|
||||
save_checkpoint(
|
||||
ckpt,
|
||||
hps.sample_rate,
|
||||
hps.if_f0,
|
||||
hps.name + "_e%s_s%s" % (epoch, global_step),
|
||||
epoch,
|
||||
hps.version,
|
||||
hps,
|
||||
model_path=hps.model_dir
|
||||
hps
|
||||
),
|
||||
)
|
||||
)
|
||||
@@ -663,7 +658,11 @@ def train_and_evaluate(
|
||||
"saving final ckpt:%s"
|
||||
% (
|
||||
save_checkpoint(
|
||||
ckpt, hps.sample_rate, hps.if_f0, hps.name, epoch, hps.version, hps
|
||||
ckpt,
|
||||
hps.name,
|
||||
epoch,
|
||||
hps,
|
||||
model_path=hps.model_path
|
||||
)
|
||||
)
|
||||
)
|
||||
@@ -673,4 +672,5 @@ def train_and_evaluate(
|
||||
|
||||
if __name__ == "__main__":
|
||||
torch.multiprocessing.set_start_method("spawn")
|
||||
main()
|
||||
hps = utils.get_hparams()
|
||||
train_model(hps)
|
||||
+30
-10
@@ -52,11 +52,8 @@ function addPreviewWidget(nodeType, nodeData, widgetName="audio", when="onNodeCr
|
||||
}
|
||||
|
||||
function previewAudio(node,options={filename: null, type: "input", widgetName: "audio", widgetId: null}){
|
||||
// while (node.widgets.length > 2){
|
||||
// node.widgets.pop();
|
||||
// }
|
||||
|
||||
try {
|
||||
console.log({node, options})
|
||||
var el = document.getElementById(options.widgetId);
|
||||
el?.remove();
|
||||
} catch (error) {
|
||||
@@ -149,19 +146,19 @@ function chainCallback(object, property, callback) {
|
||||
}
|
||||
}
|
||||
|
||||
async function uploadFile(file) {
|
||||
async function uploadFile(file,subfolder) {
|
||||
//TODO: Add uploaded file to cache with Cache.put()?
|
||||
try {
|
||||
// Wrap file in formdata so it includes filename
|
||||
const body = new FormData();
|
||||
const i = file.webkitRelativePath.lastIndexOf('/');
|
||||
const subfolder = file.webkitRelativePath.slice(0,i+1)
|
||||
// const i = file.webkitRelativePath.lastIndexOf('/');
|
||||
// const subfolder = file.webkitRelativePath.slice(0,i+1)
|
||||
const new_file = new File([file], file.name, {
|
||||
type: file.type,
|
||||
lastModified: file.lastModified,
|
||||
});
|
||||
body.append("image", new_file);
|
||||
if (i > 0) {
|
||||
if (subfolder) {
|
||||
body.append("subfolder", subfolder);
|
||||
}
|
||||
const resp = await api.fetchApi("/upload/image", {
|
||||
@@ -252,7 +249,7 @@ function addUploadWidget(nodeType, nodeData, widgetName, type="video") {
|
||||
style: "display: none",
|
||||
onchange: async () => {
|
||||
if (fileInput.files.length) {
|
||||
if (await uploadFile(fileInput.files[0]) != 200) {
|
||||
if (await uploadFile(fileInput.files[0],"audio") != 200) {
|
||||
//upload failed and file can not be added to options
|
||||
return;
|
||||
}
|
||||
@@ -266,7 +263,27 @@ function addUploadWidget(nodeType, nodeData, widgetName, type="video") {
|
||||
}
|
||||
},
|
||||
});
|
||||
}else {
|
||||
} else if (type == "zip") {
|
||||
Object.assign(fileInput, {
|
||||
type: "file",
|
||||
accept: "application/zip",
|
||||
style: "display: none",
|
||||
onchange: async () => {
|
||||
if (fileInput.files.length) {
|
||||
if (await uploadFile(fileInput.files[0],"datasets") != 200) {
|
||||
//upload failed and file can not be added to options
|
||||
return;
|
||||
}
|
||||
const filename = fileInput.files[0].name;
|
||||
pathWidget.options.values.push(filename);
|
||||
pathWidget.value = filename;
|
||||
if (pathWidget.callback) {
|
||||
pathWidget.callback(filename)
|
||||
}
|
||||
}
|
||||
},
|
||||
});
|
||||
} else {
|
||||
throw "Unknown upload type"
|
||||
}
|
||||
document.body.append(fileInput);
|
||||
@@ -289,6 +306,9 @@ app.registerExtension({
|
||||
addUploadWidget(nodeType, nodeData, "audio", "audio")
|
||||
addPreviewWidget(nodeType, nodeData, "audio", "onNodeCreated" )
|
||||
break
|
||||
case "RVCProcessDatasetNode":
|
||||
addUploadWidget(nodeType, nodeData, "dataset", "zip")
|
||||
break
|
||||
case "DownloadAudio":
|
||||
addPreviewWidget(nodeType, nodeData, "audio", "onExecuted" )
|
||||
break
|
||||
|
||||
Reference in New Issue
Block a user