Clicko777/RVC_HFv2
0
1import os2import logging3import sys4from dotenv import load_dotenv5 6load_dotenv()7 8os.environ["OMP_NUM_THREADS"] = "4"9if sys.platform == "darwin":10 os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1"11 12now_dir = os.getcwd()13sys.path.append(now_dir)14import multiprocessing15 16logger = logging.getLogger(__name__)17 18 19class Harvest(multiprocessing.Process):20 def __init__(self, inp_q, opt_q):21 multiprocessing.Process.__init__(self)22 self.inp_q = inp_q23 self.opt_q = opt_q24 25 def run(self):26 import numpy as np27 import pyworld28 29 while 1:30 idx, x, res_f0, n_cpu, ts = self.inp_q.get()31 f0, t = pyworld.harvest(32 x.astype(np.double),33 fs=16000,34 f0_ceil=1100,35 f0_floor=50,36 frame_period=10,37 )38 res_f0[idx] = f039 if len(res_f0.keys()) >= n_cpu:40 self.opt_q.put(ts)41 42 43if __name__ == "__main__":44 import json45 import multiprocessing46 import re47 import threading48 import time49 import traceback50 from multiprocessing import Queue, cpu_count51 from queue import Empty52 53 import librosa54 from tools.torchgate import TorchGate55 import numpy as np56 import PySimpleGUI as sg57 import sounddevice as sd58 import torch59 import torch.nn.functional as F60 import torchaudio.transforms as tat61 62 import tools.rvc_for_realtime as rvc_for_realtime63 from i18n.i18n import I18nAuto64 65 i18n = I18nAuto()66 device = rvc_for_realtime.config.device67 # device = torch.device(68 # "cuda"69 # if torch.cuda.is_available()70 # else ("mps" if torch.backends.mps.is_available() else "cpu")71 # )72 current_dir = os.getcwd()73 inp_q = Queue()74 opt_q = Queue()75 n_cpu = min(cpu_count(), 8)76 for _ in range(n_cpu):77 Harvest(inp_q, opt_q).start()78 79 class GUIConfig:80 def __init__(self) -> None:81 self.pth_path: str = ""82 self.index_path: str = ""83 self.pitch: int = 084 self.samplerate: int = 4000085 self.block_time: float = 1.0 # s86 self.buffer_num: int = 187 self.threhold: int = -6088 self.crossfade_time: float = 0.0489 self.extra_time: float = 2.090 self.I_noise_reduce = False91 self.O_noise_reduce = False92 self.rms_mix_rate = 0.093 self.index_rate = 0.394 self.n_cpu = min(n_cpu, 6)95 self.f0method = "harvest"96 self.sg_input_device = ""97 self.sg_output_device = ""98 99 class GUI:100 def __init__(self) -> None:101 self.config = GUIConfig()102 self.flag_vc = False103 104 self.launcher()105 106 def load(self):107 input_devices, output_devices, _, _ = self.get_devices()108 try:109 with open("configs/config.json", "r") as j:110 data = json.load(j)111 data["pm"] = data["f0method"] == "pm"112 data["harvest"] = data["f0method"] == "harvest"113 data["crepe"] = data["f0method"] == "crepe"114 data["rmvpe"] = data["f0method"] == "rmvpe"115 except:116 with open("configs/config.json", "w") as j:117 data = {118 "pth_path": " ",119 "index_path": " ",120 "sg_input_device": input_devices[sd.default.device[0]],121 "sg_output_device": output_devices[sd.default.device[1]],122 "threhold": "-60",123 "pitch": "0",124 "index_rate": "0",125 "rms_mix_rate": "0",126 "block_time": "0.25",127 "crossfade_length": "0.04",128 "extra_time": "2",129 "f0method": "rmvpe",130 }131 data["pm"] = data["f0method"] == "pm"132 data["harvest"] = data["f0method"] == "harvest"133 data["crepe"] = data["f0method"] == "crepe"134 data["rmvpe"] = data["f0method"] == "rmvpe"135 return data136 137 def launcher(self):138 data = self.load()139 sg.theme("LightBlue3")140 input_devices, output_devices, _, _ = self.get_devices()141 layout = [142 [143 sg.Frame(144 title=i18n("加载模型"),145 layout=[146 [147 sg.Input(148 default_text=data.get("pth_path", ""),149 key="pth_path",150 ),151 sg.FileBrowse(152 i18n("选择.pth文件"),153 initial_folder=os.path.join(154 os.getcwd(), "assets/weights"155 ),156 file_types=((". pth"),),157 ),158 ],159 [160 sg.Input(161 default_text=data.get("index_path", ""),162 key="index_path",163 ),164 sg.FileBrowse(165 i18n("选择.index文件"),166 initial_folder=os.path.join(os.getcwd(), "logs"),167 file_types=((". index"),),168 ),169 ],170 ],171 )172 ],173 [174 sg.Frame(175 layout=[176 [177 sg.Text(i18n("输入设备")),178 sg.Combo(179 input_devices,180 key="sg_input_device",181 default_value=data.get("sg_input_device", ""),182 ),183 ],184 [185 sg.Text(i18n("输出设备")),186 sg.Combo(187 output_devices,188 key="sg_output_device",189 default_value=data.get("sg_output_device", ""),190 ),191 ],192 [sg.Button(i18n("重载设备列表"), key="reload_devices")],193 ],194 title=i18n("音频设备(请使用同种类驱动)"),195 )196 ],197 [198 sg.Frame(199 layout=[200 [201 sg.Text(i18n("响应阈值")),202 sg.Slider(203 range=(-60, 0),204 key="threhold",205 resolution=1,206 orientation="h",207 default_value=data.get("threhold", "-60"),208 enable_events=True,209 ),210 ],211 [212 sg.Text(i18n("音调设置")),213 sg.Slider(214 range=(-24, 24),215 key="pitch",216 resolution=1,217 orientation="h",218 default_value=data.get("pitch", "0"),219 enable_events=True,220 ),221 ],222 [223 sg.Text(i18n("Index Rate")),224 sg.Slider(225 range=(0.0, 1.0),226 key="index_rate",227 resolution=0.01,228 orientation="h",229 default_value=data.get("index_rate", "0"),230 enable_events=True,231 ),232 ],233 [234 sg.Text(i18n("响度因子")),235 sg.Slider(236 range=(0.0, 1.0),237 key="rms_mix_rate",238 resolution=0.01,239 orientation="h",240 default_value=data.get("rms_mix_rate", "0"),241 enable_events=True,242 ),243 ],244 [245 sg.Text(i18n("音高算法")),246 sg.Radio(247 "pm",248 "f0method",249 key="pm",250 default=data.get("pm", "") == True,251 enable_events=True,252 ),253 sg.Radio(254 "harvest",255 "f0method",256 key="harvest",257 default=data.get("harvest", "") == True,258 enable_events=True,259 ),260 sg.Radio(261 "crepe",262 "f0method",263 key="crepe",264 default=data.get("crepe", "") == True,265 enable_events=True,266 ),267 sg.Radio(268 "rmvpe",269 "f0method",270 key="rmvpe",271 default=data.get("rmvpe", "") == True,272 enable_events=True,273 ),274 ],275 ],276 title=i18n("常规设置"),277 ),278 sg.Frame(279 layout=[280 [281 sg.Text(i18n("采样长度")),282 sg.Slider(283 range=(0.05, 2.4),284 key="block_time",285 resolution=0.01,286 orientation="h",287 default_value=data.get("block_time", "0.25"),288 enable_events=True,289 ),290 ],291 [292 sg.Text(i18n("harvest进程数")),293 sg.Slider(294 range=(1, n_cpu),295 key="n_cpu",296 resolution=1,297 orientation="h",298 default_value=data.get(299 "n_cpu", min(self.config.n_cpu, n_cpu)300 ),301 enable_events=True,302 ),303 ],304 [305 sg.Text(i18n("淡入淡出长度")),306 sg.Slider(307 range=(0.01, 0.15),308 key="crossfade_length",309 resolution=0.01,310 orientation="h",311 default_value=data.get("crossfade_length", "0.04"),312 enable_events=True,313 ),314 ],315 [316 sg.Text(i18n("额外推理时长")),317 sg.Slider(318 range=(0.05, 5.00),319 key="extra_time",320 resolution=0.01,321 orientation="h",322 default_value=data.get("extra_time", "2.0"),323 enable_events=True,324 ),325 ],326 [327 sg.Checkbox(328 i18n("输入降噪"),329 key="I_noise_reduce",330 enable_events=True,331 ),332 sg.Checkbox(333 i18n("输出降噪"),334 key="O_noise_reduce",335 enable_events=True,336 ),337 ],338 ],339 title=i18n("性能设置"),340 ),341 ],342 [343 sg.Button(i18n("开始音频转换"), key="start_vc"),344 sg.Button(i18n("停止音频转换"), key="stop_vc"),345 sg.Text(i18n("推理时间(ms):")),346 sg.Text("0", key="infer_time"),347 ],348 ]349 self.window = sg.Window("RVC - GUI", layout=layout, finalize=True)350 self.event_handler()351 352 def event_handler(self):353 while True:354 event, values = self.window.read()355 if event == sg.WINDOW_CLOSED:356 self.flag_vc = False357 exit()358 if event == "reload_devices":359 prev_input = self.window["sg_input_device"].get()360 prev_output = self.window["sg_output_device"].get()361 input_devices, output_devices, _, _ = self.get_devices(update=True)362 if prev_input not in input_devices:363 self.config.sg_input_device = input_devices[0]364 else:365 self.config.sg_input_device = prev_input366 self.window["sg_input_device"].Update(values=input_devices)367 self.window["sg_input_device"].Update(368 value=self.config.sg_input_device369 )370 if prev_output not in output_devices:371 self.config.sg_output_device = output_devices[0]372 else:373 self.config.sg_output_device = prev_output374 self.window["sg_output_device"].Update(values=output_devices)375 self.window["sg_output_device"].Update(376 value=self.config.sg_output_device377 )378 if event == "start_vc" and self.flag_vc == False:379 if self.set_values(values) == True:380 logger.info("Use CUDA: %s", torch.cuda.is_available())381 self.start_vc()382 settings = {383 "pth_path": values["pth_path"],384 "index_path": values["index_path"],385 "sg_input_device": values["sg_input_device"],386 "sg_output_device": values["sg_output_device"],387 "threhold": values["threhold"],388 "pitch": values["pitch"],389 "rms_mix_rate": values["rms_mix_rate"],390 "index_rate": values["index_rate"],391 "block_time": values["block_time"],392 "crossfade_length": values["crossfade_length"],393 "extra_time": values["extra_time"],394 "n_cpu": values["n_cpu"],395 "f0method": ["pm", "harvest", "crepe", "rmvpe"][396 [397 values["pm"],398 values["harvest"],399 values["crepe"],400 values["rmvpe"],401 ].index(True)402 ],403 }404 with open("configs/config.json", "w") as j:405 json.dump(settings, j)406 if event == "stop_vc" and self.flag_vc == True:407 self.flag_vc = False408 409 # Parameter hot update410 if event == "threhold":411 self.config.threhold = values["threhold"]412 elif event == "pitch":413 self.config.pitch = values["pitch"]414 if hasattr(self, "rvc"):415 self.rvc.change_key(values["pitch"])416 elif event == "index_rate":417 self.config.index_rate = values["index_rate"]418 if hasattr(self, "rvc"):419 self.rvc.change_index_rate(values["index_rate"])420 elif event == "rms_mix_rate":421 self.config.rms_mix_rate = values["rms_mix_rate"]422 elif event in ["pm", "harvest", "crepe", "rmvpe"]:423 self.config.f0method = event424 elif event == "I_noise_reduce":425 self.config.I_noise_reduce = values["I_noise_reduce"]426 elif event == "O_noise_reduce":427 self.config.O_noise_reduce = values["O_noise_reduce"]428 elif event != "start_vc" and self.flag_vc == True:429 # Other parameters do not support hot update430 self.flag_vc = False431 432 def set_values(self, values):433 if len(values["pth_path"].strip()) == 0:434 sg.popup(i18n("请选择pth文件"))435 return False436 if len(values["index_path"].strip()) == 0:437 sg.popup(i18n("请选择index文件"))438 return False439 pattern = re.compile("[^\x00-\x7F]+")440 if pattern.findall(values["pth_path"]):441 sg.popup(i18n("pth文件路径不可包含中文"))442 return False443 if pattern.findall(values["index_path"]):444 sg.popup(i18n("index文件路径不可包含中文"))445 return False446 self.set_devices(values["sg_input_device"], values["sg_output_device"])447 self.config.pth_path = values["pth_path"]448 self.config.index_path = values["index_path"]449 self.config.threhold = values["threhold"]450 self.config.pitch = values["pitch"]451 self.config.block_time = values["block_time"]452 self.config.crossfade_time = values["crossfade_length"]453 self.config.extra_time = values["extra_time"]454 self.config.I_noise_reduce = values["I_noise_reduce"]455 self.config.O_noise_reduce = values["O_noise_reduce"]456 self.config.rms_mix_rate = values["rms_mix_rate"]457 self.config.index_rate = values["index_rate"]458 self.config.n_cpu = values["n_cpu"]459 self.config.f0method = ["pm", "harvest", "crepe", "rmvpe"][460 [461 values["pm"],462 values["harvest"],463 values["crepe"],464 values["rmvpe"],465 ].index(True)466 ]467 return True468 469 def start_vc(self):470 torch.cuda.empty_cache()471 self.flag_vc = True472 self.rvc = rvc_for_realtime.RVC(473 self.config.pitch,474 self.config.pth_path,475 self.config.index_path,476 self.config.index_rate,477 self.config.n_cpu,478 inp_q,479 opt_q,480 device,481 self.rvc if hasattr(self, "rvc") else None482 )483 self.config.samplerate = self.rvc.tgt_sr484 self.zc = self.rvc.tgt_sr // 100485 self.block_frame = int(np.round(self.config.block_time * self.config.samplerate / self.zc)) * self.zc486 self.block_frame_16k = 160 * self.block_frame // self.zc487 self.crossfade_frame = int(np.round(self.config.crossfade_time * self.config.samplerate / self.zc)) * self.zc488 self.sola_search_frame = self.zc489 self.extra_frame = int(np.round(self.config.extra_time * self.config.samplerate / self.zc)) * self.zc490 self.input_wav: torch.Tensor = torch.zeros(491 self.extra_frame492 + self.crossfade_frame493 + self.sola_search_frame494 + self.block_frame,495 device=device,496 dtype=torch.float32,497 )498 self.input_wav_res: torch.Tensor= torch.zeros(160 * self.input_wav.shape[0] // self.zc, device=device,dtype=torch.float32)499 self.pitch: np.ndarray = np.zeros(500 self.input_wav.shape[0] // self.zc,501 dtype="int32",502 )503 self.pitchf: np.ndarray = np.zeros(504 self.input_wav.shape[0] // self.zc,505 dtype="float64",506 )507 self.sola_buffer: torch.Tensor = torch.zeros(508 self.crossfade_frame, device=device, dtype=torch.float32509 )510 self.nr_buffer: torch.Tensor = self.sola_buffer.clone()511 self.output_buffer: torch.Tensor = self.input_wav.clone()512 self.res_buffer: torch.Tensor = torch.zeros(2 * self.zc, device=device,dtype=torch.float32)513 self.valid_rate = 1 - (self.extra_frame - 1) / self.input_wav.shape[0]514 self.fade_in_window: torch.Tensor = (515 torch.sin(516 0.5517 * np.pi518 * torch.linspace(519 0.0,520 1.0,521 steps=self.crossfade_frame,522 device=device,523 dtype=torch.float32,524 )525 )526 ** 2527 )528 self.fade_out_window: torch.Tensor = 1 - self.fade_in_window529 self.resampler = tat.Resample(530 orig_freq=self.config.samplerate, new_freq=16000, dtype=torch.float32531 ).to(device)532 self.tg = TorchGate(sr=self.config.samplerate, n_fft=4*self.zc, prop_decrease=0.9).to(device)533 thread_vc = threading.Thread(target=self.soundinput)534 thread_vc.start()535 536 def soundinput(self):537 """538 接受音频输入539 """540 channels = 1 if sys.platform == "darwin" else 2541 with sd.Stream(542 channels=channels,543 callback=self.audio_callback,544 blocksize=self.block_frame,545 samplerate=self.config.samplerate,546 dtype="float32",547 ):548 while self.flag_vc:549 time.sleep(self.config.block_time)550 logger.debug("Audio block passed.")551 logger.debug("ENDing VC")552 553 def audio_callback(554 self, indata: np.ndarray, outdata: np.ndarray, frames, times, status555 ):556 """557 音频处理558 """559 start_time = time.perf_counter()560 indata = librosa.to_mono(indata.T)561 if self.config.threhold > -60:562 rms = librosa.feature.rms(563 y=indata, frame_length=4*self.zc, hop_length=self.zc564 )565 db_threhold = (566 librosa.amplitude_to_db(rms, ref=1.0)[0] < self.config.threhold567 )568 for i in range(db_threhold.shape[0]):569 if db_threhold[i]:570 indata[i * self.zc : (i + 1) * self.zc] = 0571 self.input_wav[: -self.block_frame] = self.input_wav[self.block_frame :].clone()572 self.input_wav[-self.block_frame: ] = torch.from_numpy(indata).to(device)573 self.input_wav_res[ : -self.block_frame_16k] = self.input_wav_res[self.block_frame_16k :].clone()574 # input noise reduction and resampling575 if self.config.I_noise_reduce:576 input_wav = self.input_wav[-self.crossfade_frame -self.block_frame-2*self.zc: ]577 input_wav = self.tg(input_wav.unsqueeze(0), self.input_wav.unsqueeze(0))[0, 2*self.zc:]578 input_wav[: self.crossfade_frame] *= self.fade_in_window579 input_wav[: self.crossfade_frame] += self.nr_buffer * self.fade_out_window580 self.nr_buffer[:] = input_wav[-self.crossfade_frame: ]581 input_wav = torch.cat((self.res_buffer[:], input_wav[: self.block_frame]))582 self.res_buffer[:] = input_wav[-2*self.zc: ]583 self.input_wav_res[-self.block_frame_16k-160: ] = self.resampler(input_wav)[160: ]584 else:585 self.input_wav_res[-self.block_frame_16k-160: ] = self.resampler(self.input_wav[-self.block_frame-2*self.zc: ])[160: ]586 # infer587 f0_extractor_frame = self.block_frame_16k + 800588 if self.config.f0method == 'rmvpe':589 f0_extractor_frame = 5120 * ((f0_extractor_frame - 1) // 5120 + 1)590 infer_wav = self.rvc.infer(591 self.input_wav_res,592 self.input_wav_res[-f0_extractor_frame :].cpu().numpy(),593 self.block_frame_16k,594 self.valid_rate,595 self.pitch,596 self.pitchf,597 self.config.f0method,598 )599 infer_wav = infer_wav[600 -self.crossfade_frame - self.sola_search_frame - self.block_frame :601 ]602 # output noise reduction603 if self.config.O_noise_reduce:604 self.output_buffer[: -self.block_frame] = self.output_buffer[self.block_frame :].clone()605 self.output_buffer[-self.block_frame: ] = infer_wav[-self.block_frame:]606 infer_wav = self.tg(infer_wav.unsqueeze(0), self.output_buffer.unsqueeze(0)).squeeze(0)607 # volume envelop mixing608 if self.config.rms_mix_rate < 1:609 rms1 = librosa.feature.rms(610 y=self.input_wav_res[-160*infer_wav.shape[0]//self.zc :].cpu().numpy(), 611 frame_length=640, 612 hop_length=160,613 )614 rms1 = torch.from_numpy(rms1).to(device)615 rms1 = F.interpolate(616 rms1.unsqueeze(0), size=infer_wav.shape[0] + 1, mode="linear",align_corners=True,617 )[0,0,:-1]618 rms2 = librosa.feature.rms(619 y=infer_wav[:].cpu().numpy(), frame_length=4*self.zc, hop_length=self.zc620 )621 rms2 = torch.from_numpy(rms2).to(device)622 rms2 = F.interpolate(623 rms2.unsqueeze(0), size=infer_wav.shape[0] + 1, mode="linear",align_corners=True,624 )[0,0,:-1]625 rms2 = torch.max(rms2, torch.zeros_like(rms2) + 1e-3)626 infer_wav *= torch.pow(rms1 / rms2, torch.tensor(1 - self.config.rms_mix_rate))627 # SOLA algorithm from https://github.com/yxlllc/DDSP-SVC628 conv_input = infer_wav[None, None, : self.crossfade_frame + self.sola_search_frame]629 cor_nom = F.conv1d(conv_input, self.sola_buffer[None, None, :])630 cor_den = torch.sqrt(631 F.conv1d(conv_input ** 2, torch.ones(1, 1, self.crossfade_frame, device=device)) + 1e-8)632 if sys.platform == "darwin":633 _, sola_offset = torch.max(cor_nom[0, 0] / cor_den[0, 0])634 sola_offset = sola_offset.item()635 else:636 sola_offset = torch.argmax(cor_nom[0, 0] / cor_den[0, 0])637 logger.debug("sola_offset = %d", int(sola_offset))638 infer_wav = infer_wav[sola_offset: sola_offset + self.block_frame + self.crossfade_frame]639 infer_wav[: self.crossfade_frame] *= self.fade_in_window640 infer_wav[: self.crossfade_frame] += self.sola_buffer *self.fade_out_window641 self.sola_buffer[:] = infer_wav[-self.crossfade_frame:]642 if sys.platform == "darwin":643 outdata[:] = infer_wav[:-self.crossfade_frame].cpu().numpy()[:, np.newaxis]644 else:645 outdata[:] = infer_wav[:-self.crossfade_frame].repeat(2, 1).t().cpu().numpy()646 total_time = time.perf_counter() - start_time647 self.window["infer_time"].update(int(total_time * 1000))648 logger.info("Infer time: %.2f", total_time)649 650 def get_devices(self, update: bool = True):651 """获取设备列表"""652 if update:653 sd._terminate()654 sd._initialize()655 devices = sd.query_devices()656 hostapis = sd.query_hostapis()657 for hostapi in hostapis:658 for device_idx in hostapi["devices"]:659 devices[device_idx]["hostapi_name"] = hostapi["name"]660 input_devices = [661 f"{d['name']} ({d['hostapi_name']})"662 for d in devices663 if d["max_input_channels"] > 0664 ]665 output_devices = [666 f"{d['name']} ({d['hostapi_name']})"667 for d in devices668 if d["max_output_channels"] > 0669 ]670 input_devices_indices = [671 d["index"] if "index" in d else d["name"]672 for d in devices673 if d["max_input_channels"] > 0674 ]675 output_devices_indices = [676 d["index"] if "index" in d else d["name"]677 for d in devices678 if d["max_output_channels"] > 0679 ]680 return (681 input_devices,682 output_devices,683 input_devices_indices,684 output_devices_indices,685 )686 687 def set_devices(self, input_device, output_device):688 """设置输出设备"""689 (690 input_devices,691 output_devices,692 input_device_indices,693 output_device_indices,694 ) = self.get_devices()695 sd.default.device[0] = input_device_indices[696 input_devices.index(input_device)697 ]698 sd.default.device[1] = output_device_indices[699 output_devices.index(output_device)700 ]701 logger.info(702 "Input device: %s:%s", str(sd.default.device[0]), input_device703 )704 logger.info(705 "Output device: %s:%s", str(sd.default.device[1]), output_device706 )707 708 gui = GUI()