MLVLab/Human_Object_Interaction
1
1import argparse2import datetime3import json4import random5import time6import multiprocessing7from pathlib import Path8import os9import cv210import numpy as np11import torch12from torch.utils.data import DataLoader, DistributedSampler13import hotr.data.datasets as datasets14import hotr.util.misc as utils15from hotr.engine.arg_parser import get_args_parser16from hotr.data.datasets import build_dataset, get_coco_api_from_dataset17from hotr.data.datasets.vcoco import make_hoi_transforms18from PIL import Image19from hotr.util.logger import print_params, print_args20 21import copy22from hotr.data.datasets import builtin_meta23from PIL import Image24import requests25# import mmcv26from matplotlib import pyplot as plt27import imageio28 29from tools.vis_tool import *30from hotr.models.detr import build31 32def change_format(results,valid_ids):33 34 boxes,labels,pair_score =\35 list(map(lambda x: x.cpu().numpy(), [results['boxes'], results['labels'], results['pair_score']]))36 output_i={}37 output_i['predictions']=[]38 output_i['hoi_prediction']=[]39 40 h_idx=np.where(labels==1)[0]41 for box,label in zip(boxes,labels):42 43 output_i['predictions'].append({'bbox':box.tolist(),'category_id':label})44 45 for i,verb in enumerate(pair_score):46 if i in [1,4,10,23,26,5,18]:47 continue48 for j,hum in enumerate(h_idx):49 for k in range(len(boxes)):50 if verb[j][k]>0:51 output_i['hoi_prediction'].append({'subject_id':hum,'object_id':k,'category_id':i+2,'score':verb[j][k]})52 53 return output_i54def vis(args,input_img=None,id=294,return_img=False):55 56 if args.frozen_weights is not None:57 print("Freeze weights for detector")58 if not torch.cuda.is_available():59 args.device = 'cpu'60 device = torch.device(args.device)61 62 # fix the seed for reproducibility63 seed = args.seed + utils.get_rank()64 torch.manual_seed(seed)65 np.random.seed(seed)66 random.seed(seed)67 68 # Data Setup69 dataset_train = build_dataset(image_set='train', args=args)70 args.num_classes = dataset_train.num_category()71 args.num_actions = dataset_train.num_action()72 args.action_names = dataset_train.get_actions()73 if args.share_enc: args.hoi_enc_layers = args.enc_layers74 if args.pretrained_dec: args.hoi_dec_layers = args.dec_layers75 if args.dataset_file == 'vcoco':76 # Save V-COCO dataset statistics77 args.valid_ids = np.array(dataset_train.get_object_label_idx()).nonzero()[0]78 args.invalid_ids = np.argwhere(np.array(dataset_train.get_object_label_idx()) == 0).squeeze(1)79 args.human_actions = dataset_train.get_human_action()80 args.object_actions = dataset_train.get_object_action()81 args.num_human_act = dataset_train.num_human_act()82 elif args.dataset_file == 'hico-det':83 args.valid_obj_ids = dataset_train.get_valid_obj_ids()84 print_args(args)85 86 args.HOIDet=True87 args.eval=True88 args.pretrained_dec=True89 args.share_enc=True90 args.share_dec_param = True91 if args.dataset_file=='hico-det':92 args.valid_ids=args.valid_obj_ids93 94 # Model Setup95 model, criterion, postprocessors = build(args)96 model.to(device)97 98 model_without_ddp = model99 100 n_parameters = print_params(model)101 102 param_dicts = [103 {"params": [p for n, p in model_without_ddp.named_parameters() if "backbone" not in n and p.requires_grad]},104 {105 "params": [p for n, p in model_without_ddp.named_parameters() if "backbone" in n and p.requires_grad],106 "lr": args.lr_backbone,107 },108 ]109 110 output_dir = Path(args.output_dir)111 112 checkpoint = torch.load(args.resume, map_location='cpu')113 #수정114 module_name=list(checkpoint['model'].keys())115 model_without_ddp.load_state_dict(checkpoint['model'], strict=False)116 117 # if not args.video_vis:118 # url='http://images.cocodataset.org/val2014/COCO_val2014_{}.jpg'.format(str(id).zfill(12))119 # req = requests.get(url, stream=True, timeout=1, verify=False).raw120 121 if input_img is None:122 req = args.image_dir123 img = Image.open(req).convert('RGB')124 else:125 # import pdb;pdb.set_trace()126 img = input_img127 128 w,h=img.size129 orig_size = torch.as_tensor([int(h), int(w)]).unsqueeze(0).to(device)130 131 transform=make_hoi_transforms('val')132 sample=img.copy()133 sample,_=transform(sample,None)134 sample = sample.unsqueeze(0).to(device)135 with torch.no_grad():136 model.eval()137 out=model(sample)138 results = postprocessors['hoi'](out, orig_size,dataset=args.dataset_file,args=args)139 output_i=change_format(results[0],args.valid_ids)140 141 out_dir = './vis'142 image = np.asarray(img, dtype=np.uint8)[:,:,::-1]143 # image = cv2.imdecode(image_nparray, cv2.IMREAD_COLOR)144 145 vis_img=draw_img_vcoco(image,output_i,top_k=args.topk,threshold=args.threshold,color=builtin_meta.COCO_CATEGORIES) 146 plt.imshow(cv2.cvtColor(vis_img,cv2.COLOR_BGR2RGB))147 148 if return_img:149 vis_img = Image.fromarray(vis_img[:,:,::-1])150 # import pdb;pdb.set_trace()151 return vis_img152 else:153 cv2.imwrite('./vis_res/vis1.jpg',vis_img)154 155 # else:156 # frames=[]157 # video_file=id158 159 # video_reader = mmcv.VideoReader('./vid/'+video_file+'.mp4')160 # fourcc = cv2.VideoWriter_fourcc(*'mp4v')161 # video_writer = cv2.VideoWriter(162 # './vid/'+video_file+'_vis.mp4', fourcc, video_reader.fps,163 # (video_reader.width, video_reader.height))164 165 # orig_size = torch.as_tensor([int(video_reader.height), int(video_reader.width)]).unsqueeze(0).to(device)166 # transform=make_hoi_transforms('val')167 168 # for frame in mmcv.track_iter_progress(video_reader):169 170 # frame=mmcv.imread(frame)171 # frame=frame.copy()172 173 # frame=Image.fromarray(frame,'RGB')174 175 # sample,_=transform(frame,None)176 # sample=sample.unsqueeze(0).to(device)177 178 # with torch.no_grad():179 # model.eval()180 # out=model(sample)181 # results = postprocessors['hoi'](out, orig_size,dataset='vcoco',args=args)182 # output_i=change_format(results[0],args.valid_ids)183 184 # vis_img=draw_img_vcoco(np.array(frame),output_i,top_k=args.topk,threshold=args.threshold,color=builtin_meta.COCO_CATEGORIES)185 # frames.append(vis_img)186 # video_writer.write(vis_img)187 188 # with imageio.get_writer("smiling.gif", mode="I") as writer:189 # for idx, frame in enumerate(frames):190 # # print("Adding frame to GIF file: ", idx + 1)191 # writer.append_data(frame)192 # if video_writer:193 # video_writer.release()194 # cv2.destroyAllWindows()195 196 197# def visualization(id, video_vis=False, dataset_file='vcoco', path_id = 0 ,data_path='v-coco', threshold=0.4, topk=10,aug_path = '[]'):198 199# parser = argparse.ArgumentParser('DETR training and evaluation script', parents=[get_args_parser()])200# checkpoint_dir= './checkpoints/vcoco/checkpoint.pth' if dataset_file=='vcoco' else './checkpoints/hico-det/hico_ft_q16.pth'201# with open('./v-coco/data/vcoco_test.ids') as file:202# test_idxs = [line.rstrip('\n') for line in file]203# if not video_vis:204# id = test_idxs[id]205# args = parser.parse_args(args=['--dataset_file',dataset_file,'--data_path',data_path,'--resume',checkpoint_dir,'--num_hoi_queries' ,'16','--temperature' ,'0.05', '--augpath_name',aug_path ,'--path_id','{}'.format(path_id)])206# args.video_vis=video_vis207# args.threshold=threshold208# args.topk=topk209 210# if args.output_dir:211# Path(args.output_dir).mkdir(parents=True, exist_ok=True)212# vis(args,id)213 214# 230727 for huggingface215def visualization(input_img,threshold,topk):216 217 parser = argparse.ArgumentParser('DETR training and evaluation script', parents=[get_args_parser()])218 args = parser.parse_args(args=[])219 args.threshold = threshold220 args.topk = int(topk)221 222 # checkpoint_dir= './checkpoints/vcoco/checkpoint.pth' if dataset_file=='vcoco' else './checkpoints/hico-det/hico_ft_q16.pth'223 args.resume= './checkpoints/vcoco/checkpoint.pth'224 # with open('./v-coco/data/splits/vcoco_test.ids') as file:225 # test_idxs = [line.rstrip('\n') for line in file]226 # # if not video_vis:227 # id = test_idxs[309]228 # args = parser.parse_args()229 args.dataset_file = 'vcoco'230 args.data_path = 'v-coco'231 # args.resume = checkpoint_dir232 args.num_hoi_queries = 16233 args.temperature = 0.05234 args.augpath_name = ['p2','p3','p4']235 # args.path_id = 1236 # args.threshold = threshold237 # args.topk = topk238 if args.output_dir:239 Path(args.output_dir).mkdir(parents=True, exist_ok=True)240 return vis(args,input_img=input_img,return_img=True)241 242if __name__ == '__main__':243 parser = argparse.ArgumentParser('DETR training and evaluation script', parents=[get_args_parser()])244 parser.add_argument('--threshold',help='score threshold for visualization', default=0.4, type=float)245 # parser.add_argument('--path_id',help='index of inference path', default=1, type=int)246 parser.add_argument('--topk',help='topk prediction', default=5, type=int)247 parser.add_argument('--video_vis', action='store_true')248 parser.add_argument('--image_dir', default='', type=str)249 args = parser.parse_args()250 # checkpoint_dir= './checkpoints/vcoco/checkpoint.pth' if dataset_file=='vcoco' else './checkpoints/hico-det/hico_ft_q16.pth'251 args.resume= './checkpoints/vcoco/checkpoint.pth'252 with open('./v-coco/data/splits/vcoco_test.ids') as file:253 test_idxs = [line.rstrip('\n') for line in file]254 # if not video_vis:255 id = test_idxs[309]256 # args = parser.parse_args()257 # args.dataset_file = 'vcoco'258 # args.data_path = 'v-coco'259 # args.resume = checkpoint_dir260 # args.num_hoi_queries = 16261 # args.temperature = 0.05262 args.augpath_name = ['p2','p3','p4']263 # args.path_id = 1264 265 if args.output_dir:266 Path(args.output_dir).mkdir(parents=True, exist_ok=True)267 vis(args,id)268 