CoolFace
Apppublic

iti/HandMesh

sourceHugging Faceupdated 5y agoView on Hugging Face
0likes
main.py126 linesDownload Raw Back to root
1import os.path as osp2#import torch3#import torch.backends.cudnn as cudnn4from cmr.cmr_sg import CMR_SG5from cmr.cmr_pg import CMR_PG6from cmr.cmr_g import CMR_G7from mobrecon.mobrecon_densestack import MobRecon8from utils.read import spiral_tramsform9from utils import utils, writer10from options.base_options import BaseOptions11from datasets.FreiHAND.freihand import FreiHAND12from datasets.Human36M.human36m import Human36M13#from torch.utils.data import DataLoader14from run import Runner15from termcolor import cprint16from tensorboardX import SummaryWriter17 18if __name__ == '__main__':19    # get config20    args = BaseOptions().parse()21 22    # dir prepare23    args.work_dir = osp.dirname(osp.realpath(__file__))24    data_fp = osp.join(args.work_dir, 'data', args.dataset)25    args.out_dir = osp.join(args.work_dir, 'out', args.dataset, args.exp_name)26    args.checkpoints_dir = osp.join(args.out_dir, 'checkpoints')27    if args.phase in ['eval', 'demo']:28        utils.makedirs(osp.join(args.out_dir, args.phase))29    utils.makedirs(args.out_dir)30    utils.makedirs(args.checkpoints_dir)31 32    # device set33    if -1 in args.device_idx or not torch.cuda.is_available():34        device = torch.device('cpu')35    elif len(args.device_idx) == 1:36        device = torch.device('cuda', args.device_idx[0])37    else:38        device = torch.device('cuda')39    torch.set_num_threads(args.n_threads)40 41    # deterministic42    cudnn.benchmark = True43    cudnn.deterministic = True44 45    if args.dataset=='Human36M':46        template_fp = osp.join(args.work_dir, 'template', 'template_body.ply')47        transform_fp = osp.join(args.work_dir, 'template', 'transform_body.pkl')48    else:49        template_fp = osp.join(args.work_dir, 'template', 'template.ply')50        transform_fp = osp.join(args.work_dir, 'template', 'transform.pkl')51    spiral_indices_list, down_transform_list, up_transform_list, tmp = spiral_tramsform(transform_fp, template_fp, args.ds_factors, args.seq_length, args.dilation)52 53    # model54    if args.model == 'cmr_sg':55        model = CMR_SG(args, spiral_indices_list, up_transform_list)56    elif args.model == 'cmr_pg':57        model = CMR_PG(args, spiral_indices_list, up_transform_list)58    elif args.model == 'cmr_g':59        model = CMR_G(args, spiral_indices_list, up_transform_list)60    elif args.model == 'mobrecon':61        for i in range(len(up_transform_list)):62            up_transform_list[i] = (*up_transform_list[i]._indices(), up_transform_list[i]._values())63        model = MobRecon(args, spiral_indices_list, up_transform_list)64    else:65        raise Exception('Model {} not support'.format(args.model))66 67    # load68    epoch = 069    if args.resume:70        if len(args.resume.split('/')) > 1:71            model_path = args.resume72        else:73            model_path = osp.join(args.checkpoints_dir, args.resume)74        checkpoint = torch.load(model_path, map_location='cpu')75        if checkpoint.get('model_state_dict', None) is not None:76            checkpoint = checkpoint['model_state_dict']77        model.load_state_dict(checkpoint)78        epoch = checkpoint.get('epoch', -1) + 179        cprint('Load checkpoint {}'.format(model_path), 'yellow')80    model = model.to(device)81 82    # run83    runner = Runner(args, model, tmp['face'], device)84    if args.phase == 'train':85        # log86        writer = writer.Writer(args)87        writer.print_str(args)88        # dataset89        if args.dataset=='FreiHAND':90            eval_dataset = FreiHAND(data_fp, 'evaluation', args, tmp['face'])91            eval_loader = DataLoader(eval_dataset, batch_size=1, shuffle=False, pin_memory=True, num_workers=0)92            train_dataset = FreiHAND(data_fp, 'training', args, tmp['face'], writer=writer, down_sample_list=down_transform_list, ms=args.ms_mesh)93            train_loader = DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True, pin_memory=False, num_workers=16, drop_last=True)94        elif args.dataset=='Human36M':95            eval_dataset = Human36M(data_fp, 'test', args, down_transform_list, tmp['face'])96            eval_loader = DataLoader(eval_dataset, batch_size=1, shuffle=False, pin_memory=True, num_workers=0)97            train_dataset = Human36M(data_fp, 'train', args, down_transform_list, tmp['face'])98            train_loader = DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True, pin_memory=False, num_workers=16, drop_last=True)99        else:100            raise Exception('Dataset not support')101        # optimize102        optimizer = torch.optim.Adam(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)103        scheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer, args.decay_step, gamma=args.lr_decay)104        # tensorboard105        board = SummaryWriter(osp.join(args.out_dir, 'board'))106        runner.set_train_loader(train_loader, args.epochs, optimizer, scheduler, writer, board, start_epoch=epoch)107        runner.set_eval_loader(eval_loader)108        runner.train()109    elif args.phase == 'eval':110        # dataset111        eval_dataset = FreiHAND(data_fp, 'evaluation', args, tmp['face'])112        eval_loader = DataLoader(eval_dataset, batch_size=1, shuffle=False, pin_memory=True, num_workers=0)113        runner.set_eval_loader(eval_loader)114        runner.evaluation()115    elif args.phase == 'eval_withgt':116        # dataset117        eval_dataset = Human36M(data_fp, 'test', args, down_transform_list, tmp['face'])118        eval_loader = DataLoader(eval_dataset, batch_size=1, shuffle=False, pin_memory=True, num_workers=0)119        runner.set_eval_loader(eval_loader)120        runner.evaluation_withgt()121    elif args.phase == 'demo':122        runner.set_demo(args)123        runner.demo()124    else:125        raise Exception('phase error')126