iti/HandMesh
0
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 