Team Ai
Apppublic

MLVLab/Human_Object_Interaction

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
1likes
visualization.py268 linesDownload Raw Back to root
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