goathead777/Zero_Shot_Inference
0
1import os2import glob3import sys4import argparse5import logging6import json7import subprocess8import traceback9 10import librosa11import numpy as np12from scipy.io.wavfile import read13import torch14import logging15 16logging.getLogger("numba").setLevel(logging.ERROR)17logging.getLogger("matplotlib").setLevel(logging.ERROR)18 19MATPLOTLIB_FLAG = False20 21# logging.basicConfig(stream=sys.stdout, level=logging.DEBUG)22logger = logging23 24 25def load_checkpoint(checkpoint_path, model, optimizer=None, skip_optimizer=False):26 assert os.path.isfile(checkpoint_path)27 checkpoint_dict = torch.load(checkpoint_path, map_location="cpu")28 iteration = checkpoint_dict["iteration"]29 learning_rate = checkpoint_dict["learning_rate"]30 if (31 optimizer is not None32 and not skip_optimizer33 and checkpoint_dict["optimizer"] is not None34 ):35 optimizer.load_state_dict(checkpoint_dict["optimizer"])36 saved_state_dict = checkpoint_dict["model"]37 if hasattr(model, "module"):38 state_dict = model.module.state_dict()39 else:40 state_dict = model.state_dict()41 new_state_dict = {}42 for k, v in state_dict.items():43 try:44 # assert "quantizer" not in k45 # print("load", k)46 new_state_dict[k] = saved_state_dict[k]47 assert saved_state_dict[k].shape == v.shape, (48 saved_state_dict[k].shape,49 v.shape,50 )51 except:52 traceback.print_exc()53 print(54 "error, %s is not in the checkpoint" % k55 ) # shape不对也会,比如text_embedding当cleaner修改时56 new_state_dict[k] = v57 if hasattr(model, "module"):58 model.module.load_state_dict(new_state_dict)59 else:60 model.load_state_dict(new_state_dict)61 print("load ")62 logger.info(63 "Loaded checkpoint '{}' (iteration {})".format(checkpoint_path, iteration)64 )65 return model, optimizer, learning_rate, iteration66 67 68def save_checkpoint(model, optimizer, learning_rate, iteration, checkpoint_path):69 logger.info(70 "Saving model and optimizer state at iteration {} to {}".format(71 iteration, checkpoint_path72 )73 )74 if hasattr(model, "module"):75 state_dict = model.module.state_dict()76 else:77 state_dict = model.state_dict()78 torch.save(79 {80 "model": state_dict,81 "iteration": iteration,82 "optimizer": optimizer.state_dict(),83 "learning_rate": learning_rate,84 },85 checkpoint_path,86 )87 88 89def summarize(90 writer,91 global_step,92 scalars={},93 histograms={},94 images={},95 audios={},96 audio_sampling_rate=22050,97):98 for k, v in scalars.items():99 writer.add_scalar(k, v, global_step)100 for k, v in histograms.items():101 writer.add_histogram(k, v, global_step)102 for k, v in images.items():103 writer.add_image(k, v, global_step, dataformats="HWC")104 for k, v in audios.items():105 writer.add_audio(k, v, global_step, audio_sampling_rate)106 107 108def latest_checkpoint_path(dir_path, regex="G_*.pth"):109 f_list = glob.glob(os.path.join(dir_path, regex))110 f_list.sort(key=lambda f: int("".join(filter(str.isdigit, f))))111 x = f_list[-1]112 print(x)113 return x114 115 116def plot_spectrogram_to_numpy(spectrogram):117 global MATPLOTLIB_FLAG118 if not MATPLOTLIB_FLAG:119 import matplotlib120 121 matplotlib.use("Agg")122 MATPLOTLIB_FLAG = True123 mpl_logger = logging.getLogger("matplotlib")124 mpl_logger.setLevel(logging.WARNING)125 import matplotlib.pylab as plt126 import numpy as np127 128 fig, ax = plt.subplots(figsize=(10, 2))129 im = ax.imshow(spectrogram, aspect="auto", origin="lower", interpolation="none")130 plt.colorbar(im, ax=ax)131 plt.xlabel("Frames")132 plt.ylabel("Channels")133 plt.tight_layout()134 135 fig.canvas.draw()136 data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep="")137 data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,))138 plt.close()139 return data140 141 142def plot_alignment_to_numpy(alignment, info=None):143 global MATPLOTLIB_FLAG144 if not MATPLOTLIB_FLAG:145 import matplotlib146 147 matplotlib.use("Agg")148 MATPLOTLIB_FLAG = True149 mpl_logger = logging.getLogger("matplotlib")150 mpl_logger.setLevel(logging.WARNING)151 import matplotlib.pylab as plt152 import numpy as np153 154 fig, ax = plt.subplots(figsize=(6, 4))155 im = ax.imshow(156 alignment.transpose(), aspect="auto", origin="lower", interpolation="none"157 )158 fig.colorbar(im, ax=ax)159 xlabel = "Decoder timestep"160 if info is not None:161 xlabel += "\n\n" + info162 plt.xlabel(xlabel)163 plt.ylabel("Encoder timestep")164 plt.tight_layout()165 166 fig.canvas.draw()167 data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep="")168 data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,))169 plt.close()170 return data171 172 173def load_wav_to_torch(full_path):174 data, sampling_rate = librosa.load(full_path, sr=None)175 return torch.FloatTensor(data), sampling_rate176 177 178def load_filepaths_and_text(filename, split="|"):179 with open(filename, encoding="utf-8") as f:180 filepaths_and_text = [line.strip().split(split) for line in f]181 return filepaths_and_text182 183 184def get_hparams(init=True, stage=1):185 parser = argparse.ArgumentParser()186 parser.add_argument(187 "-c",188 "--config",189 type=str,190 default="./configs/s2.json",191 help="JSON file for configuration",192 )193 parser.add_argument(194 "-p", "--pretrain", type=str, required=False, default=None, help="pretrain dir"195 )196 parser.add_argument(197 "-rs",198 "--resume_step",199 type=int,200 required=False,201 default=None,202 help="resume step",203 )204 # parser.add_argument('-e', '--exp_dir', type=str, required=False,default=None,help='experiment directory')205 # parser.add_argument('-g', '--pretrained_s2G', type=str, required=False,default=None,help='pretrained sovits gererator weights')206 # parser.add_argument('-d', '--pretrained_s2D', type=str, required=False,default=None,help='pretrained sovits discriminator weights')207 208 args = parser.parse_args()209 210 config_path = args.config211 with open(config_path, "r") as f:212 data = f.read()213 config = json.loads(data)214 215 hparams = HParams(**config)216 hparams.pretrain = args.pretrain217 hparams.resume_step = args.resume_step218 # hparams.data.exp_dir = args.exp_dir219 if stage == 1:220 model_dir = hparams.s1_ckpt_dir221 else:222 model_dir = hparams.s2_ckpt_dir223 config_save_path = os.path.join(model_dir, "config.json")224 225 if not os.path.exists(model_dir):226 os.makedirs(model_dir)227 228 with open(config_save_path, "w") as f:229 f.write(data)230 return hparams231 232 233def clean_checkpoints(path_to_models="logs/44k/", n_ckpts_to_keep=2, sort_by_time=True):234 """Freeing up space by deleting saved ckpts235 236 Arguments:237 path_to_models -- Path to the model directory238 n_ckpts_to_keep -- Number of ckpts to keep, excluding G_0.pth and D_0.pth239 sort_by_time -- True -> chronologically delete ckpts240 False -> lexicographically delete ckpts241 """242 import re243 244 ckpts_files = [245 f246 for f in os.listdir(path_to_models)247 if os.path.isfile(os.path.join(path_to_models, f))248 ]249 name_key = lambda _f: int(re.compile("._(\d+)\.pth").match(_f).group(1))250 time_key = lambda _f: os.path.getmtime(os.path.join(path_to_models, _f))251 sort_key = time_key if sort_by_time else name_key252 x_sorted = lambda _x: sorted(253 [f for f in ckpts_files if f.startswith(_x) and not f.endswith("_0.pth")],254 key=sort_key,255 )256 to_del = [257 os.path.join(path_to_models, fn)258 for fn in (x_sorted("G")[:-n_ckpts_to_keep] + x_sorted("D")[:-n_ckpts_to_keep])259 ]260 del_info = lambda fn: logger.info(f".. Free up space by deleting ckpt {fn}")261 del_routine = lambda x: [os.remove(x), del_info(x)]262 rs = [del_routine(fn) for fn in to_del]263 264 265def get_hparams_from_dir(model_dir):266 config_save_path = os.path.join(model_dir, "config.json")267 with open(config_save_path, "r") as f:268 data = f.read()269 config = json.loads(data)270 271 hparams = HParams(**config)272 hparams.model_dir = model_dir273 return hparams274 275 276def get_hparams_from_file(config_path):277 with open(config_path, "r") as f:278 data = f.read()279 config = json.loads(data)280 281 hparams = HParams(**config)282 return hparams283 284 285def check_git_hash(model_dir):286 source_dir = os.path.dirname(os.path.realpath(__file__))287 if not os.path.exists(os.path.join(source_dir, ".git")):288 logger.warn(289 "{} is not a git repository, therefore hash value comparison will be ignored.".format(290 source_dir291 )292 )293 return294 295 cur_hash = subprocess.getoutput("git rev-parse HEAD")296 297 path = os.path.join(model_dir, "githash")298 if os.path.exists(path):299 saved_hash = open(path).read()300 if saved_hash != cur_hash:301 logger.warn(302 "git hash values are different. {}(saved) != {}(current)".format(303 saved_hash[:8], cur_hash[:8]304 )305 )306 else:307 open(path, "w").write(cur_hash)308 309 310def get_logger(model_dir, filename="train.log"):311 global logger312 logger = logging.getLogger(os.path.basename(model_dir))313 logger.setLevel(logging.DEBUG)314 315 formatter = logging.Formatter("%(asctime)s\t%(name)s\t%(levelname)s\t%(message)s")316 if not os.path.exists(model_dir):317 os.makedirs(model_dir)318 h = logging.FileHandler(os.path.join(model_dir, filename))319 h.setLevel(logging.DEBUG)320 h.setFormatter(formatter)321 logger.addHandler(h)322 return logger323 324 325class HParams:326 def __init__(self, **kwargs):327 for k, v in kwargs.items():328 if type(v) == dict:329 v = HParams(**v)330 self[k] = v331 332 def keys(self):333 return self.__dict__.keys()334 335 def items(self):336 return self.__dict__.items()337 338 def values(self):339 return self.__dict__.values()340 341 def __len__(self):342 return len(self.__dict__)343 344 def __getitem__(self, key):345 return getattr(self, key)346 347 def __setitem__(self, key, value):348 return setattr(self, key, value)349 350 def __contains__(self, key):351 return key in self.__dict__352 353 def __repr__(self):354 return self.__dict__.__repr__()355 356 357if __name__ == "__main__":358 print(359 load_wav_to_torch(360 "/home/fish/wenetspeech/dataset_vq/Y0000022499_wHFSeHEx9CM/S00261.flac"361 )362 )363 