CoolFace
Apppublic

Paolify/RVC_4

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
gui_v1.py708 linesDownload Raw Back to root
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()