MLVLab/Human_Object_Interaction
1
1# ------------------------------------------------------------------------2# HOTR official code : main.py3# Copyright (c) Kakao Brain, Inc. and its affiliates. All Rights Reserved4# ------------------------------------------------------------------------5# Modified from DETR (https://github.com/facebookresearch/detr)6# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved7# ------------------------------------------------------------------------8import argparse9import datetime10import json11import random12import time13import multiprocessing14from pathlib import Path15 16import numpy as np17import torch18from torch.utils.data import DataLoader, DistributedSampler19 20import hotr.data.datasets as datasets21import hotr.util.misc as utils22from hotr.engine.arg_parser import get_args_parser23from hotr.data.datasets import build_dataset, get_coco_api_from_dataset24from hotr.engine.trainer import train_one_epoch25from hotr.engine import hoi_evaluator, hoi_accumulator26from hotr.models import build_model27import wandb28 29from hotr.util.logger import print_params, print_args30 31def save_ckpt(args, model_without_ddp, optimizer, lr_scheduler, epoch, filename):32 # save_ckpt: function for saving checkpoints33 output_dir = Path(args.output_dir)34 if args.output_dir:35 checkpoint_path = output_dir / f'{filename}.pth'36 utils.save_on_master({37 'model': model_without_ddp.state_dict(),38 'optimizer': optimizer.state_dict(),39 'lr_scheduler': lr_scheduler.state_dict(),40 'epoch': epoch,41 'args': args,42 }, checkpoint_path)43 44def main(args):45 utils.init_distributed_mode(args)46 47 if args.frozen_weights is not None:48 print("Freeze weights for detector")49 50 if not torch.cuda.is_available():51 args.device = 'cpu'52 device = torch.device(args.device)53 54 # fix the seed for reproducibility55 seed = args.seed + utils.get_rank()56 torch.manual_seed(seed)57 np.random.seed(seed)58 random.seed(seed)59 60 # Data Setup61 dataset_train = build_dataset(image_set='train', args=args)62 dataset_val = build_dataset(image_set='val' if not args.eval else 'test', args=args)63 assert dataset_train.num_action() == dataset_val.num_action(), "Number of actions should be the same between splits"64 args.num_classes = dataset_train.num_category()65 args.num_actions = dataset_train.num_action()66 args.action_names = dataset_train.get_actions()67 if args.share_enc: args.hoi_enc_layers = args.enc_layers68 if args.pretrained_dec: args.hoi_dec_layers = args.dec_layers69 if args.dataset_file == 'vcoco':70 # Save V-COCO dataset statistics71 args.valid_ids = np.array(dataset_train.get_object_label_idx()).nonzero()[0]72 args.invalid_ids = np.argwhere(np.array(dataset_train.get_object_label_idx()) == 0).squeeze(1)73 args.human_actions = dataset_train.get_human_action()74 args.object_actions = dataset_train.get_object_action()75 args.num_human_act = dataset_train.num_human_act()76 elif args.dataset_file == 'hico-det':77 args.valid_obj_ids = dataset_train.get_valid_obj_ids()78 print_args(args)79 80 if args.distributed:81 sampler_train = DistributedSampler(dataset_train, shuffle=True)82 sampler_val = DistributedSampler(dataset_val, shuffle=False)83 else:84 sampler_train = torch.utils.data.RandomSampler(dataset_train)85 sampler_val = torch.utils.data.SequentialSampler(dataset_val)86 87 batch_sampler_train = torch.utils.data.BatchSampler(88 sampler_train, args.batch_size, drop_last=True)89 90 data_loader_train = DataLoader(dataset_train, batch_sampler=batch_sampler_train,91 collate_fn=utils.collate_fn, num_workers=args.num_workers)92 data_loader_val = DataLoader(dataset_val, args.batch_size, sampler=sampler_val,93 drop_last=False, collate_fn=utils.collate_fn, num_workers=args.num_workers)94 95 # Model Setup96 model, criterion, postprocessors = build_model(args)97 # import pdb;pdb.set_trace()98 model.to(device)99 100 model_without_ddp = model101 if args.distributed:102 model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[args.gpu])103 model_without_ddp = model.module104 n_parameters = print_params(model)105 106 param_dicts = [107 {"params": [p for n, p in model_without_ddp.named_parameters() if "backbone" not in n and p.requires_grad]},108 {109 "params": [p for n, p in model_without_ddp.named_parameters() if "backbone" in n and p.requires_grad],110 "lr": args.lr_backbone,111 },112 ]113 optimizer = torch.optim.AdamW(param_dicts, lr=args.lr, weight_decay=args.weight_decay)114 lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, args.lr_drop)115 # lr_scheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer, [1,100])116 117 118 # Weight Setup119 if args.frozen_weights is not None:120 if args.frozen_weights.startswith('https'):121 checkpoint = torch.hub.load_state_dict_from_url(122 args.frozen_weights, map_location='cpu', check_hash=True)123 else:124 checkpoint = torch.load(args.frozen_weights, map_location='cpu')125 model_without_ddp.detr.load_state_dict(checkpoint['model'])126 127 if args.resume:128 if args.resume.startswith('https'):129 checkpoint = torch.hub.load_state_dict_from_url(130 args.resume, map_location='cpu', check_hash=True)131 else:132 checkpoint = torch.load(args.resume, map_location='cpu')133 model_without_ddp.load_state_dict(checkpoint['model'])134 if not args.eval and 'optimizer' in checkpoint and 'lr_scheduler' in checkpoint and 'epoch' in checkpoint:135 optimizer.load_state_dict(checkpoint['optimizer'])136 # lr_scheduler.load_state_dict(checkpoint['lr_scheduler'])137 args.start_epoch = checkpoint['epoch'] + 1138 # import pdb;pdb.set_trace()139 if args.eval:140 # test only mode141 if args.HOIDet:142 if args.dataset_file == 'vcoco':143 total_res = hoi_evaluator(args, model, criterion, postprocessors, data_loader_val, device)144 sc1, sc2 = hoi_accumulator(args, total_res, True, False)145 elif args.dataset_file == 'hico-det':146 test_stats = hoi_evaluator(args, model, None, postprocessors, data_loader_val, device)147 print(f'| mAP (full)\t\t: {test_stats["mAP"]:.2f}')148 print(f'| mAP (rare)\t\t: {test_stats["mAP rare"]:.2f}')149 print(f'| mAP (non-rare)\t: {test_stats["mAP non-rare"]:.2f}')150 else: raise ValueError(f'dataset {args.dataset_file} is not supported.')151 return152 else:153 test_stats, coco_evaluator = evaluate_coco(model, criterion, postprocessors,154 data_loader_val, base_ds, device, args.output_dir)155 if args.output_dir:156 utils.save_on_master(coco_evaluator.coco_eval["bbox"].eval, output_dir / "eval.pth")157 return158 159 # stats160 scenario1, scenario2 = 0, 0161 best_mAP, best_rare, best_non_rare = 0, 0, 0162 163 # add argparse164 if args.wandb and utils.get_rank() == 0:165 wandb.init(166 project=args.project_name,167 group=args.group_name,168 name=args.run_name,169 config=args170 )171 wandb.watch(model)172 173 # Training starts here!174 # lr_scheduler.step()175 start_time = time.time()176 for epoch in range(args.start_epoch, args.epochs):177 if args.distributed:178 sampler_train.set_epoch(epoch)179 train_stats = train_one_epoch(180 model, criterion, data_loader_train, optimizer, device, epoch, args.epochs, args.ramp_up_epoch,args.ramp_down_epoch,args.hoi_consistency_loss_coef,181 args.clip_max_norm, dataset_file=args.dataset_file, log=args.wandb)182 lr_scheduler.step()183 184 # Validation185 if args.validate:186 print('-'*100)187 if args.dataset_file == 'vcoco':188 total_res = hoi_evaluator(args, model, criterion, postprocessors, data_loader_val, device)189 if utils.get_rank() == 0:190 sc1, sc2 = hoi_accumulator(args, total_res, False, args.wandb)191 if sc1 > scenario1:192 scenario1 = sc1193 scenario2 = sc2194 save_ckpt(args, model_without_ddp, optimizer, lr_scheduler, epoch, filename='best')195 print(f'| Scenario #1 mAP : {sc1:.2f} ({scenario1:.2f})')196 print(f'| Scenario #2 mAP : {sc2:.2f} ({scenario2:.2f})')197 elif args.dataset_file == 'hico-det':198 test_stats = hoi_evaluator(args, model, None, postprocessors, data_loader_val, device)199 if utils.get_rank() == 0:200 if test_stats['mAP'] > best_mAP:201 best_mAP = test_stats['mAP']202 best_rare = test_stats['mAP rare']203 best_non_rare = test_stats['mAP non-rare']204 save_ckpt(args, model_without_ddp, optimizer, lr_scheduler, epoch, filename='best')205 print(f'| mAP (full)\t\t: {test_stats["mAP"]:.2f} ({best_mAP:.2f})')206 print(f'| mAP (rare)\t\t: {test_stats["mAP rare"]:.2f} ({best_rare:.2f})')207 print(f'| mAP (non-rare)\t: {test_stats["mAP non-rare"]:.2f} ({best_non_rare:.2f})')208 if args.wandb and utils.get_rank() == 0:209 wandb.log({210 'mAP': test_stats['mAP'],211 'mAP rare': test_stats['mAP rare'],212 'mAP non-rare': test_stats['mAP non-rare']213 })214 print('-'*100)215 216 save_ckpt(args, model_without_ddp, optimizer, lr_scheduler, epoch, filename='checkpoint')217 if (epoch + 1) % args.lr_drop == 0 :218 save_ckpt(args, model_without_ddp, optimizer, lr_scheduler, epoch, filename='checkpoint_'+str(epoch))219 # if (epoch + 1) % args.pseudo_epoch == 0 :220 # save_ckpt(args, model_without_ddp, optimizer, lr_scheduler, epoch, filename='checkpoint_pseudo_'+str(epoch))221 total_time = time.time() - start_time222 total_time_str = str(datetime.timedelta(seconds=int(total_time)))223 print('Training time {}'.format(total_time_str))224 if args.dataset_file == 'vcoco':225 print(f'| Scenario #1 mAP : {scenario1:.2f}')226 print(f'| Scenario #2 mAP : {scenario2:.2f}')227 elif args.dataset_file == 'hico-det':228 print(f'| mAP (full)\t\t: {best_mAP:.2f}')229 print(f'| mAP (rare)\t\t: {best_rare:.2f}')230 print(f'| mAP (non-rare)\t: {best_non_rare:.2f}')231 232 233if __name__ == '__main__':234 parser = argparse.ArgumentParser(235 'End-to-End Human Object Interaction training and evaluation script',236 parents=[get_args_parser()]237 )238 args = parser.parse_args()239 if args.output_dir:240 args.output_dir += f"/{args.group_name}/{args.run_name}/"241 Path(args.output_dir).mkdir(parents=True, exist_ok=True)242 main(args)243 