CoolFace
Apppublic

ChazzyG/Retrieval-based-Voice-Conversion-WebUI

sourceHugging Faceapache-2.0updated 9mo agoView on Hugging Face
0likes
utils.py472 linesDownload Raw Back to train
1import os, traceback2import glob3import sys4import argparse5import logging6import json7import subprocess8import numpy as np9from scipy.io.wavfile import read10import torch11 12MATPLOTLIB_FLAG = False13 14logging.basicConfig(stream=sys.stdout, level=logging.DEBUG)15logger = logging16 17 18def load_checkpoint_d(checkpoint_path, combd, sbd, optimizer=None, load_opt=1):19    assert os.path.isfile(checkpoint_path)20    checkpoint_dict = torch.load(checkpoint_path, map_location="cpu")21 22    ##################23    def go(model, bkey):24        saved_state_dict = checkpoint_dict[bkey]25        if hasattr(model, "module"):26            state_dict = model.module.state_dict()27        else:28            state_dict = model.state_dict()29        new_state_dict = {}30        for k, v in state_dict.items():  # 模型需要的shape31            try:32                new_state_dict[k] = saved_state_dict[k]33                if saved_state_dict[k].shape != state_dict[k].shape:34                    print(35                        "shape-%s-mismatch|need-%s|get-%s"36                        % (k, state_dict[k].shape, saved_state_dict[k].shape)37                    )  #38                    raise KeyError39            except:40                # logger.info(traceback.format_exc())41                logger.info("%s is not in the checkpoint" % k)  # pretrain缺失的42                new_state_dict[k] = v  # 模型自带的随机值43        if hasattr(model, "module"):44            model.module.load_state_dict(new_state_dict, strict=False)45        else:46            model.load_state_dict(new_state_dict, strict=False)47 48    go(combd, "combd")49    go(sbd, "sbd")50    #############51    logger.info("Loaded model weights")52 53    iteration = checkpoint_dict["iteration"]54    learning_rate = checkpoint_dict["learning_rate"]55    if (56        optimizer is not None and load_opt == 157    ):  ###加载不了,如果是空的的话,重新初始化,可能还会影响lr时间表的更新,因此在train文件最外围catch58        #   try:59        optimizer.load_state_dict(checkpoint_dict["optimizer"])60    #   except:61    #     traceback.print_exc()62    logger.info("Loaded checkpoint '{}' (epoch {})".format(checkpoint_path, iteration))63    return model, optimizer, learning_rate, iteration64 65 66# def load_checkpoint(checkpoint_path, model, optimizer=None):67#   assert os.path.isfile(checkpoint_path)68#   checkpoint_dict = torch.load(checkpoint_path, map_location='cpu')69#   iteration = checkpoint_dict['iteration']70#   learning_rate = checkpoint_dict['learning_rate']71#   if optimizer is not None:72#     optimizer.load_state_dict(checkpoint_dict['optimizer'])73#   # print(1111)74#   saved_state_dict = checkpoint_dict['model']75#   # print(1111)76#77#   if hasattr(model, 'module'):78#     state_dict = model.module.state_dict()79#   else:80#     state_dict = model.state_dict()81#   new_state_dict= {}82#   for k, v in state_dict.items():83#     try:84#       new_state_dict[k] = saved_state_dict[k]85#     except:86#       logger.info("%s is not in the checkpoint" % k)87#       new_state_dict[k] = v88#   if hasattr(model, 'module'):89#     model.module.load_state_dict(new_state_dict)90#   else:91#     model.load_state_dict(new_state_dict)92#   logger.info("Loaded checkpoint '{}' (epoch {})" .format(93#     checkpoint_path, iteration))94#   return model, optimizer, learning_rate, iteration95def load_checkpoint(checkpoint_path, model, optimizer=None, load_opt=1):96    assert os.path.isfile(checkpoint_path)97    checkpoint_dict = torch.load(checkpoint_path, map_location="cpu")98 99    saved_state_dict = checkpoint_dict["model"]100    if hasattr(model, "module"):101        state_dict = model.module.state_dict()102    else:103        state_dict = model.state_dict()104    new_state_dict = {}105    for k, v in state_dict.items():  # 模型需要的shape106        try:107            new_state_dict[k] = saved_state_dict[k]108            if saved_state_dict[k].shape != state_dict[k].shape:109                print(110                    "shape-%s-mismatch|need-%s|get-%s"111                    % (k, state_dict[k].shape, saved_state_dict[k].shape)112                )  #113                raise KeyError114        except:115            # logger.info(traceback.format_exc())116            logger.info("%s is not in the checkpoint" % k)  # pretrain缺失的117            new_state_dict[k] = v  # 模型自带的随机值118    if hasattr(model, "module"):119        model.module.load_state_dict(new_state_dict, strict=False)120    else:121        model.load_state_dict(new_state_dict, strict=False)122    logger.info("Loaded model weights")123 124    iteration = checkpoint_dict["iteration"]125    learning_rate = checkpoint_dict["learning_rate"]126    if (127        optimizer is not None and load_opt == 1128    ):  ###加载不了,如果是空的的话,重新初始化,可能还会影响lr时间表的更新,因此在train文件最外围catch129        #   try:130        optimizer.load_state_dict(checkpoint_dict["optimizer"])131    #   except:132    #     traceback.print_exc()133    logger.info("Loaded checkpoint '{}' (epoch {})".format(checkpoint_path, iteration))134    return model, optimizer, learning_rate, iteration135 136 137def save_checkpoint(model, optimizer, learning_rate, iteration, checkpoint_path):138    logger.info(139        "Saving model and optimizer state at epoch {} to {}".format(140            iteration, checkpoint_path141        )142    )143    if hasattr(model, "module"):144        state_dict = model.module.state_dict()145    else:146        state_dict = model.state_dict()147    torch.save(148        {149            "model": state_dict,150            "iteration": iteration,151            "optimizer": optimizer.state_dict(),152            "learning_rate": learning_rate,153        },154        checkpoint_path,155    )156 157 158def save_checkpoint_d(combd, sbd, optimizer, learning_rate, iteration, checkpoint_path):159    logger.info(160        "Saving model and optimizer state at epoch {} to {}".format(161            iteration, checkpoint_path162        )163    )164    if hasattr(combd, "module"):165        state_dict_combd = combd.module.state_dict()166    else:167        state_dict_combd = combd.state_dict()168    if hasattr(sbd, "module"):169        state_dict_sbd = sbd.module.state_dict()170    else:171        state_dict_sbd = sbd.state_dict()172    torch.save(173        {174            "combd": state_dict_combd,175            "sbd": state_dict_sbd,176            "iteration": iteration,177            "optimizer": optimizer.state_dict(),178            "learning_rate": learning_rate,179        },180        checkpoint_path,181    )182 183 184def summarize(185    writer,186    global_step,187    scalars={},188    histograms={},189    images={},190    audios={},191    audio_sampling_rate=22050,192):193    for k, v in scalars.items():194        writer.add_scalar(k, v, global_step)195    for k, v in histograms.items():196        writer.add_histogram(k, v, global_step)197    for k, v in images.items():198        writer.add_image(k, v, global_step, dataformats="HWC")199    for k, v in audios.items():200        writer.add_audio(k, v, global_step, audio_sampling_rate)201 202 203def latest_checkpoint_path(dir_path, regex="G_*.pth"):204    f_list = glob.glob(os.path.join(dir_path, regex))205    f_list.sort(key=lambda f: int("".join(filter(str.isdigit, f))))206    x = f_list[-1]207    print(x)208    return x209 210 211def plot_spectrogram_to_numpy(spectrogram):212    global MATPLOTLIB_FLAG213    if not MATPLOTLIB_FLAG:214        import matplotlib215 216        matplotlib.use("Agg")217        MATPLOTLIB_FLAG = True218        mpl_logger = logging.getLogger("matplotlib")219        mpl_logger.setLevel(logging.WARNING)220    import matplotlib.pylab as plt221    import numpy as np222 223    fig, ax = plt.subplots(figsize=(10, 2))224    im = ax.imshow(spectrogram, aspect="auto", origin="lower", interpolation="none")225    plt.colorbar(im, ax=ax)226    plt.xlabel("Frames")227    plt.ylabel("Channels")228    plt.tight_layout()229 230    fig.canvas.draw()231    data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep="")232    data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,))233    plt.close()234    return data235 236 237def plot_alignment_to_numpy(alignment, info=None):238    global MATPLOTLIB_FLAG239    if not MATPLOTLIB_FLAG:240        import matplotlib241 242        matplotlib.use("Agg")243        MATPLOTLIB_FLAG = True244        mpl_logger = logging.getLogger("matplotlib")245        mpl_logger.setLevel(logging.WARNING)246    import matplotlib.pylab as plt247    import numpy as np248 249    fig, ax = plt.subplots(figsize=(6, 4))250    im = ax.imshow(251        alignment.transpose(), aspect="auto", origin="lower", interpolation="none"252    )253    fig.colorbar(im, ax=ax)254    xlabel = "Decoder timestep"255    if info is not None:256        xlabel += "\n\n" + info257    plt.xlabel(xlabel)258    plt.ylabel("Encoder timestep")259    plt.tight_layout()260 261    fig.canvas.draw()262    data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep="")263    data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,))264    plt.close()265    return data266 267 268def load_wav_to_torch(full_path):269    sampling_rate, data = read(full_path)270    return torch.FloatTensor(data.astype(np.float32)), sampling_rate271 272 273def load_filepaths_and_text(filename, split="|"):274    with open(filename, encoding="utf-8") as f:275        filepaths_and_text = [line.strip().split(split) for line in f]276    return filepaths_and_text277 278 279def get_hparams(init=True):280    """281    todo:282      结尾七人组:283        保存频率、总epoch                     done284        bs                                    done285        pretrainG、pretrainD                  done286        卡号:os.en["CUDA_VISIBLE_DEVICES"]   done287        if_latest                             todo288      模型:if_f0                             todo289      采样率:自动选择config                  done290      是否缓存数据集进GPU:if_cache_data_in_gpu done291 292      -m:293        自动决定training_files路径,改掉train_nsf_load_pretrain.py里的hps.data.training_files    done294      -c不要了295    """296    parser = argparse.ArgumentParser()297    # parser.add_argument('-c', '--config', type=str, default="configs/40k.json",help='JSON file for configuration')298    parser.add_argument(299        "-se",300        "--save_every_epoch",301        type=int,302        required=True,303        help="checkpoint save frequency (epoch)",304    )305    parser.add_argument(306        "-te", "--total_epoch", type=int, required=True, help="total_epoch"307    )308    parser.add_argument(309        "-pg", "--pretrainG", type=str, default="", help="Pretrained Discriminator path"310    )311    parser.add_argument(312        "-pd", "--pretrainD", type=str, default="", help="Pretrained Generator path"313    )314    parser.add_argument("-g", "--gpus", type=str, default="0", help="split by -")315    parser.add_argument(316        "-bs", "--batch_size", type=int, required=True, help="batch size"317    )318    parser.add_argument(319        "-e", "--experiment_dir", type=str, required=True, help="experiment dir"320    )  # -m321    parser.add_argument(322        "-sr", "--sample_rate", type=str, required=True, help="sample rate, 32k/40k/48k"323    )324    parser.add_argument(325        "-f0",326        "--if_f0",327        type=int,328        required=True,329        help="use f0 as one of the inputs of the model, 1 or 0",330    )331    parser.add_argument(332        "-l",333        "--if_latest",334        type=int,335        required=True,336        help="if only save the latest G/D pth file, 1 or 0",337    )338    parser.add_argument(339        "-c",340        "--if_cache_data_in_gpu",341        type=int,342        required=True,343        help="if caching the dataset in GPU memory, 1 or 0",344    )345 346    args = parser.parse_args()347    name = args.experiment_dir348    experiment_dir = os.path.join("./logs", args.experiment_dir)349 350    if not os.path.exists(experiment_dir):351        os.makedirs(experiment_dir)352 353    config_path = "configs/%s.json" % args.sample_rate354    config_save_path = os.path.join(experiment_dir, "config.json")355    if init:356        with open(config_path, "r") as f:357            data = f.read()358        with open(config_save_path, "w") as f:359            f.write(data)360    else:361        with open(config_save_path, "r") as f:362            data = f.read()363    config = json.loads(data)364 365    hparams = HParams(**config)366    hparams.model_dir = hparams.experiment_dir = experiment_dir367    hparams.save_every_epoch = args.save_every_epoch368    hparams.name = name369    hparams.total_epoch = args.total_epoch370    hparams.pretrainG = args.pretrainG371    hparams.pretrainD = args.pretrainD372    hparams.gpus = args.gpus373    hparams.train.batch_size = args.batch_size374    hparams.sample_rate = args.sample_rate375    hparams.if_f0 = args.if_f0376    hparams.if_latest = args.if_latest377    hparams.if_cache_data_in_gpu = args.if_cache_data_in_gpu378    hparams.data.training_files = "%s/filelist.txt" % experiment_dir379    return hparams380 381 382def get_hparams_from_dir(model_dir):383    config_save_path = os.path.join(model_dir, "config.json")384    with open(config_save_path, "r") as f:385        data = f.read()386    config = json.loads(data)387 388    hparams = HParams(**config)389    hparams.model_dir = model_dir390    return hparams391 392 393def get_hparams_from_file(config_path):394    with open(config_path, "r") as f:395        data = f.read()396    config = json.loads(data)397 398    hparams = HParams(**config)399    return hparams400 401 402def check_git_hash(model_dir):403    source_dir = os.path.dirname(os.path.realpath(__file__))404    if not os.path.exists(os.path.join(source_dir, ".git")):405        logger.warn(406            "{} is not a git repository, therefore hash value comparison will be ignored.".format(407                source_dir408            )409        )410        return411 412    cur_hash = subprocess.getoutput("git rev-parse HEAD")413 414    path = os.path.join(model_dir, "githash")415    if os.path.exists(path):416        saved_hash = open(path).read()417        if saved_hash != cur_hash:418            logger.warn(419                "git hash values are different. {}(saved) != {}(current)".format(420                    saved_hash[:8], cur_hash[:8]421                )422            )423    else:424        open(path, "w").write(cur_hash)425 426 427def get_logger(model_dir, filename="train.log"):428    global logger429    logger = logging.getLogger(os.path.basename(model_dir))430    logger.setLevel(logging.DEBUG)431 432    formatter = logging.Formatter("%(asctime)s\t%(name)s\t%(levelname)s\t%(message)s")433    if not os.path.exists(model_dir):434        os.makedirs(model_dir)435    h = logging.FileHandler(os.path.join(model_dir, filename))436    h.setLevel(logging.DEBUG)437    h.setFormatter(formatter)438    logger.addHandler(h)439    return logger440 441 442class HParams:443    def __init__(self, **kwargs):444        for k, v in kwargs.items():445            if type(v) == dict:446                v = HParams(**v)447            self[k] = v448 449    def keys(self):450        return self.__dict__.keys()451 452    def items(self):453        return self.__dict__.items()454 455    def values(self):456        return self.__dict__.values()457 458    def __len__(self):459        return len(self.__dict__)460 461    def __getitem__(self, key):462        return getattr(self, key)463 464    def __setitem__(self, key, value):465        return setattr(self, key, value)466 467    def __contains__(self, key):468        return key in self.__dict__469 470    def __repr__(self):471        return self.__dict__.__repr__()472