ChazzyG/Retrieval-based-Voice-Conversion-WebUI
0
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 