kwau/sovits-isla
0
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 