CoolFace
Apppublic

kwau/sovits-isla

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
train_diff.py78 linesDownload Raw Back to root
1import argparse2 3import torch4from loguru import logger5from torch.optim import lr_scheduler6 7from diffusion.data_loaders import get_data_loaders8from diffusion.logger import utils9from diffusion.solver import train10from diffusion.unit2mel import Unit2Mel11from diffusion.vocoder import Vocoder12 13 14def parse_args(args=None, namespace=None):15    """Parse command-line arguments."""16    parser = argparse.ArgumentParser()17    parser.add_argument(18        "-c",19        "--config",20        type=str,21        required=True,22        help="path to the config file")23    return parser.parse_args(args=args, namespace=namespace)24 25 26if __name__ == '__main__':27    # parse commands28    cmd = parse_args()29    30    # load config31    args = utils.load_config(cmd.config)32    logger.info(' > config:'+ cmd.config)33    logger.info(' > exp:'+ args.env.expdir)34    35    # load vocoder36    vocoder = Vocoder(args.vocoder.type, args.vocoder.ckpt, device=args.device)37    38    # load model39    model = Unit2Mel(40                args.data.encoder_out_channels, 41                args.model.n_spk,42                args.model.use_pitch_aug,43                vocoder.dimension,44                args.model.n_layers,45                args.model.n_chans,46                args.model.n_hidden,47                args.model.timesteps,48                args.model.k_step_max49                )50    51    logger.info(f' > Now model timesteps is {model.timesteps}, and k_step_max is {model.k_step_max}')52    53    # load parameters54    optimizer = torch.optim.AdamW(model.parameters())55    initial_global_step, model, optimizer = utils.load_model(args.env.expdir, model, optimizer, device=args.device)56    for param_group in optimizer.param_groups:57        param_group['initial_lr'] = args.train.lr58        param_group['lr'] = args.train.lr * (args.train.gamma ** max(((initial_global_step-2)//args.train.decay_step),0) )59        param_group['weight_decay'] = args.train.weight_decay60    scheduler = lr_scheduler.StepLR(optimizer, step_size=args.train.decay_step, gamma=args.train.gamma,last_epoch=initial_global_step-2)61    62    # device63    if args.device == 'cuda':64        torch.cuda.set_device(args.env.gpu_id)65    model.to(args.device)66    67    for state in optimizer.state.values():68        for k, v in state.items():69            if torch.is_tensor(v):70                state[k] = v.to(args.device)71                    72    # datas73    loader_train, loader_valid = get_data_loaders(args, whole_audio=False)74    75    # run76    train(args, initial_global_step, model, optimizer, scheduler, vocoder, loader_train, loader_valid)77    78