radames/Text2Human-API
1
1import argparse2import logging3import os4import os.path as osp5import random6import time7 8import torch9 10from data.segm_attr_dataset import DeepFashionAttrSegmDataset11from models import create_model12from utils.logger import MessageLogger, get_root_logger, init_tb_logger13from utils.options import dict2str, dict_to_nonedict, parse14from utils.util import make_exp_dirs15 16 17def main():18 # options19 parser = argparse.ArgumentParser()20 parser.add_argument('-opt', type=str, help='Path to option YAML file.')21 args = parser.parse_args()22 opt = parse(args.opt, is_train=True)23 24 # mkdir and loggers25 make_exp_dirs(opt)26 log_file = osp.join(opt['path']['log'], f"train_{opt['name']}.log")27 logger = get_root_logger(28 logger_name='base', log_level=logging.INFO, log_file=log_file)29 logger.info(dict2str(opt))30 # initialize tensorboard logger31 tb_logger = None32 if opt['use_tb_logger'] and 'debug' not in opt['name']:33 tb_logger = init_tb_logger(log_dir='./tb_logger/' + opt['name'])34 35 # convert to NoneDict, which returns None for missing keys36 opt = dict_to_nonedict(opt)37 38 # set up data loader39 train_dataset = DeepFashionAttrSegmDataset(40 img_dir=opt['train_img_dir'],41 segm_dir=opt['segm_dir'],42 pose_dir=opt['pose_dir'],43 ann_dir=opt['train_ann_file'],44 xflip=True)45 train_loader = torch.utils.data.DataLoader(46 dataset=train_dataset,47 batch_size=opt['batch_size'],48 shuffle=True,49 num_workers=opt['num_workers'],50 persistent_workers=True,51 drop_last=True)52 logger.info(f'Number of train set: {len(train_dataset)}.')53 opt['max_iters'] = opt['num_epochs'] * len(54 train_dataset) // opt['batch_size']55 56 val_dataset = DeepFashionAttrSegmDataset(57 img_dir=opt['train_img_dir'],58 segm_dir=opt['segm_dir'],59 pose_dir=opt['pose_dir'],60 ann_dir=opt['val_ann_file'])61 val_loader = torch.utils.data.DataLoader(62 dataset=val_dataset, batch_size=opt['batch_size'], shuffle=False)63 logger.info(f'Number of val set: {len(val_dataset)}.')64 65 test_dataset = DeepFashionAttrSegmDataset(66 img_dir=opt['test_img_dir'],67 segm_dir=opt['segm_dir'],68 pose_dir=opt['pose_dir'],69 ann_dir=opt['test_ann_file'])70 test_loader = torch.utils.data.DataLoader(71 dataset=test_dataset, batch_size=opt['batch_size'], shuffle=False)72 logger.info(f'Number of test set: {len(test_dataset)}.')73 74 current_iter = 075 76 model = create_model(opt)77 78 data_time, iter_time = 0, 079 current_iter = 080 81 # create message logger (formatted outputs)82 msg_logger = MessageLogger(opt, current_iter, tb_logger)83 84 for epoch in range(opt['num_epochs']):85 lr = model.update_learning_rate(epoch, current_iter)86 87 for _, batch_data in enumerate(train_loader):88 data_time = time.time() - data_time89 90 current_iter += 191 92 model.feed_data(batch_data)93 model.optimize_parameters()94 95 iter_time = time.time() - iter_time96 if current_iter % opt['print_freq'] == 0:97 log_vars = {'epoch': epoch, 'iter': current_iter}98 log_vars.update({'lrs': [lr]})99 log_vars.update({'time': iter_time, 'data_time': data_time})100 log_vars.update(model.get_current_log())101 msg_logger(log_vars)102 103 data_time = time.time()104 iter_time = time.time()105 106 if epoch % opt['val_freq'] == 0 and epoch != 0:107 save_dir = f'{opt["path"]["visualization"]}/valset/epoch_{epoch:03d}' # noqa108 os.makedirs(save_dir, exist_ok=opt['debug'])109 model.inference(val_loader, save_dir)110 111 save_dir = f'{opt["path"]["visualization"]}/testset/epoch_{epoch:03d}' # noqa112 os.makedirs(save_dir, exist_ok=opt['debug'])113 model.inference(test_loader, save_dir)114 115 # save model116 model.save_network(117 model._denoise_fn,118 f'{opt["path"]["models"]}/sampler_epoch{epoch}.pth')119 120 121if __name__ == '__main__':122 main()123 