CoolFace
Apppublic

ORI-Muchim/BarKeYaeTTS

sourceHugging Facemitupdated 3y agoView on Hugging Face
3likes
utils.py227 linesDownload Raw Back to root
1import os2import 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.ERROR)15logger = logging16 17 18def load_checkpoint(checkpoint_path, model, optimizer=None):19    assert os.path.isfile(checkpoint_path)20    checkpoint_dict = torch.load(checkpoint_path, map_location='cpu')21    iteration = checkpoint_dict['iteration']22    learning_rate = checkpoint_dict['learning_rate']23    if optimizer is not None:24        optimizer.load_state_dict(checkpoint_dict['optimizer'])25    saved_state_dict = checkpoint_dict['model']26    if hasattr(model, 'module'):27        state_dict = model.module.state_dict()28    else:29        state_dict = model.state_dict()30    new_state_dict = {}31    for k, v in state_dict.items():32        try:33            new_state_dict[k] = saved_state_dict[k]34        except:35            logger.info("%s is not in the checkpoint" % k)36            new_state_dict[k] = v37    if hasattr(model, 'module'):38        model.module.load_state_dict(new_state_dict)39    else:40        model.load_state_dict(new_state_dict)41    logger.info("Loaded checkpoint '{}' (iteration {})".format(42        checkpoint_path, iteration))43    return model, optimizer, learning_rate, iteration44 45 46def plot_spectrogram_to_numpy(spectrogram):47    global MATPLOTLIB_FLAG48    if not MATPLOTLIB_FLAG:49        import matplotlib50        matplotlib.use("Agg")51        MATPLOTLIB_FLAG = True52        mpl_logger = logging.getLogger('matplotlib')53        mpl_logger.setLevel(logging.WARNING)54    import matplotlib.pylab as plt55    import numpy as np56 57    fig, ax = plt.subplots(figsize=(10, 2))58    im = ax.imshow(spectrogram, aspect="auto", origin="lower",59                   interpolation='none')60    plt.colorbar(im, ax=ax)61    plt.xlabel("Frames")62    plt.ylabel("Channels")63    plt.tight_layout()64 65    fig.canvas.draw()66    data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep='')67    data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,))68    plt.close()69    return data70 71 72def plot_alignment_to_numpy(alignment, info=None):73    global MATPLOTLIB_FLAG74    if not MATPLOTLIB_FLAG:75        import matplotlib76        matplotlib.use("Agg")77        MATPLOTLIB_FLAG = True78        mpl_logger = logging.getLogger('matplotlib')79        mpl_logger.setLevel(logging.WARNING)80    import matplotlib.pylab as plt81    import numpy as np82 83    fig, ax = plt.subplots(figsize=(6, 4))84    im = ax.imshow(alignment.transpose(), aspect='auto', origin='lower',85                   interpolation='none')86    fig.colorbar(im, ax=ax)87    xlabel = 'Decoder timestep'88    if info is not None:89        xlabel += '\n\n' + info90    plt.xlabel(xlabel)91    plt.ylabel('Encoder timestep')92    plt.tight_layout()93 94    fig.canvas.draw()95    data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep='')96    data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,))97    plt.close()98    return data99 100 101def load_wav_to_torch(full_path):102    sampling_rate, data = read(full_path)103    return torch.FloatTensor(data.astype(np.float32)), sampling_rate104 105 106def load_filepaths_and_text(filename, split="|"):107    with open(filename, encoding='utf-8') as f:108        filepaths_and_text = [line.strip().split(split) for line in f]109    return filepaths_and_text110 111 112def get_hparams(init=True):113    parser = argparse.ArgumentParser()114    parser.add_argument('-c', '--config', type=str, default="./configs/base.json",115                        help='JSON file for configuration')116    parser.add_argument('-m', '--model', type=str, required=True,117                        help='Model name')118 119    args = parser.parse_args()120    model_dir = os.path.join("./logs", args.model)121 122    if not os.path.exists(model_dir):123        os.makedirs(model_dir)124 125    config_path = args.config126    config_save_path = os.path.join(model_dir, "config.json")127    if init:128        with open(config_path, "r") as f:129            data = f.read()130        with open(config_save_path, "w") as f:131            f.write(data)132    else:133        with open(config_save_path, "r") as f:134            data = f.read()135    config = json.loads(data)136 137    hparams = HParams(**config)138    hparams.model_dir = model_dir139    return hparams140 141 142def get_hparams_from_dir(model_dir):143    config_save_path = os.path.join(model_dir, "config.json")144    with open(config_save_path, "r") as f:145        data = f.read()146    config = json.loads(data)147 148    hparams = HParams(**config)149    hparams.model_dir = model_dir150    return hparams151 152 153def get_hparams_from_file(config_path):154    with open(config_path, "r", encoding="utf-8") as f:155        data = f.read()156    config = json.loads(data)157 158    hparams = HParams(**config)159    return hparams160 161 162def check_git_hash(model_dir):163    source_dir = os.path.dirname(os.path.realpath(__file__))164    if not os.path.exists(os.path.join(source_dir, ".git")):165        logger.warn("{} is not a git repository, therefore hash value comparison will be ignored.".format(166            source_dir167        ))168        return169 170    cur_hash = subprocess.getoutput("git rev-parse HEAD")171 172    path = os.path.join(model_dir, "githash")173    if os.path.exists(path):174        saved_hash = open(path).read()175        if saved_hash != cur_hash:176            logger.warn("git hash values are different. {}(saved) != {}(current)".format(177                saved_hash[:8], cur_hash[:8]))178    else:179        open(path, "w").write(cur_hash)180 181 182def get_logger(model_dir, filename="train.log"):183    global logger184    logger = logging.getLogger(os.path.basename(model_dir))185    logger.setLevel(logging.DEBUG)186 187    formatter = logging.Formatter("%(asctime)s\t%(name)s\t%(levelname)s\t%(message)s")188    if not os.path.exists(model_dir):189        os.makedirs(model_dir)190    h = logging.FileHandler(os.path.join(model_dir, filename))191    h.setLevel(logging.DEBUG)192    h.setFormatter(formatter)193    logger.addHandler(h)194    return logger195 196 197class HParams():198    def __init__(self, **kwargs):199        for k, v in kwargs.items():200            if type(v) == dict:201                v = HParams(**v)202            self[k] = v203 204    def keys(self):205        return self.__dict__.keys()206 207    def items(self):208        return self.__dict__.items()209 210    def values(self):211        return self.__dict__.values()212 213    def __len__(self):214        return len(self.__dict__)215 216    def __getitem__(self, key):217        return getattr(self, key)218 219    def __setitem__(self, key, value):220        return setattr(self, key, value)221 222    def __contains__(self, key):223        return key in self.__dict__224 225    def __repr__(self):226        return self.__dict__.__repr__()227