Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
train_parsing_token.py123 linesDownload Raw Back to Text2Human
1import argparse2import logging3import os4import os.path as osp5import random6import time7 8import torch9 10from data.mask_dataset import MaskDataset11from 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 = MaskDataset(40        segm_dir=opt['segm_dir'], ann_dir=opt['train_ann_file'], xflip=True)41    train_loader = torch.utils.data.DataLoader(42        dataset=train_dataset,43        batch_size=opt['batch_size'],44        shuffle=True,45        num_workers=opt['num_workers'],46        persistent_workers=True,47        drop_last=True)48    logger.info(f'Number of train set: {len(train_dataset)}.')49    opt['max_iters'] = opt['num_epochs'] * len(50        train_dataset) // opt['batch_size']51 52    val_dataset = MaskDataset(53        segm_dir=opt['segm_dir'], ann_dir=opt['val_ann_file'])54    val_loader = torch.utils.data.DataLoader(55        dataset=val_dataset, batch_size=1, shuffle=False)56    logger.info(f'Number of val set: {len(val_dataset)}.')57 58    test_dataset = MaskDataset(59        segm_dir=opt['segm_dir'], ann_dir=opt['test_ann_file'])60    test_loader = torch.utils.data.DataLoader(61        dataset=test_dataset, batch_size=1, shuffle=False)62    logger.info(f'Number of test set: {len(test_dataset)}.')63 64    current_iter = 065    best_epoch = None66    best_loss = 10000067 68    model = create_model(opt)69 70    data_time, iter_time = 0, 071    current_iter = 072 73    # create message logger (formatted outputs)74    msg_logger = MessageLogger(opt, current_iter, tb_logger)75 76    for epoch in range(opt['num_epochs']):77        lr = model.update_learning_rate(epoch)78 79        for _, batch_data in enumerate(train_loader):80            data_time = time.time() - data_time81 82            current_iter += 183 84            model.optimize_parameters(batch_data, current_iter)85 86            iter_time = time.time() - iter_time87            if current_iter % opt['print_freq'] == 0:88                log_vars = {'epoch': epoch, 'iter': current_iter}89                log_vars.update({'lrs': [lr]})90                log_vars.update({'time': iter_time, 'data_time': data_time})91                log_vars.update(model.get_current_log())92                msg_logger(log_vars)93 94            data_time = time.time()95            iter_time = time.time()96 97        if epoch % opt['val_freq'] == 0:98            save_dir = f'{opt["path"]["visualization"]}/valset/epoch_{epoch:03d}'  # noqa99            os.makedirs(save_dir, exist_ok=opt['debug'])100            val_loss_total, _, _ = model.inference(val_loader, save_dir)101 102            save_dir = f'{opt["path"]["visualization"]}/testset/epoch_{epoch:03d}'  # noqa103            os.makedirs(save_dir, exist_ok=opt['debug'])104            test_loss_total, _, _ = model.inference(test_loader, save_dir)105 106            logger.info(f'Epoch: {epoch}, '107                        f'val_loss_total: {val_loss_total}, '108                        f'test_loss_total: {test_loss_total}.')109 110            if test_loss_total < best_loss:111                best_epoch = epoch112                best_loss = test_loss_total113 114            logger.info(f'Best epoch: {best_epoch}, '115                        f'Best test loss: {best_loss: .4f}.')116 117            # save model118            model.save_network(f'{opt["path"]["models"]}/epoch{epoch}.pth')119 120 121if __name__ == '__main__':122    main()123