CVPR/GroupViT
9
1# Modified from the implementation of https://huggingface.co/akhaliq2import os3import sys4os.system("git clone https://github.com/NVlabs/GroupViT")5sys.path.insert(0, 'GroupViT')6 7import os.path as osp8from collections import namedtuple9 10import gradio as gr11import mmcv12import numpy as np13import torch14from datasets import build_text_transform15from mmcv.cnn.utils import revert_sync_batchnorm16from mmcv.image import tensor2imgs17from mmcv.parallel import collate, scatter18from models import build_model19from omegaconf import read_write20from segmentation.datasets import (COCOObjectDataset, PascalContextDataset,21 PascalVOCDataset)22from segmentation.evaluation import (GROUP_PALETTE, build_seg_demo_pipeline,23 build_seg_inference)24from utils import get_config, load_checkpoint25 26import shutil27 28if not osp.exists('GroupViT/hg_demo'):29 shutil.copytree('demo/', 'GroupViT/hg_demo/')30 31os.chdir('GroupViT')32# checkpoint_url = 'https://github.com/xvjiarui/GroupViT-1/releases/download/v1.0.0/group_vit_gcc_yfcc_30e-74d335e6.pth'33checkpoint_url = 'https://github.com/xvjiarui/GroupViT/releases/download/v1.0.0/group_vit_gcc_yfcc_30e-879422e0.pth'34cfg_path = 'configs/group_vit_gcc_yfcc_30e.yml'35output_dir = 'demo/output'36device = 'cpu'37# vis_modes = ['first_group', 'final_group', 'input_pred_label']38vis_modes = ['input_pred_label', 'final_group']39output_labels = ['segmentation map', 'groups']40dataset_options = ['Pascal VOC', 'Pascal Context', 'COCO']41examples = [['Pascal VOC', '', 'hg_demo/voc.jpg'],42 ['Pascal Context', '', 'hg_demo/ctx.jpg'],43 ['COCO', '', 'hg_demo/coco.jpg']]44 45PSEUDO_ARGS = namedtuple('PSEUDO_ARGS',46 ['cfg', 'opts', 'resume', 'vis', 'local_rank'])47 48args = PSEUDO_ARGS(49 cfg=cfg_path, opts=[], resume=checkpoint_url, vis=vis_modes, local_rank=0)50 51cfg = get_config(args)52 53with read_write(cfg):54 cfg.evaluate.eval_only = True55 56model = build_model(cfg.model)57model = revert_sync_batchnorm(model)58model.to(device)59model.eval()60 61load_checkpoint(cfg, model, None, None)62 63text_transform = build_text_transform(False, cfg.data.text_aug, with_dc=False)64test_pipeline = build_seg_demo_pipeline()65 66 67def inference(dataset, additional_classes, input_img):68 if dataset == 'voc' or dataset == 'Pascal VOC':69 dataset_class = PascalVOCDataset70 seg_cfg = 'segmentation/configs/_base_/datasets/pascal_voc12.py'71 elif dataset == 'coco' or dataset == 'COCO':72 dataset_class = COCOObjectDataset73 seg_cfg = 'segmentation/configs/_base_/datasets/coco.py'74 elif dataset == 'context' or dataset == 'Pascal Context':75 dataset_class = PascalContextDataset76 seg_cfg = 'segmentation/configs/_base_/datasets/pascal_context.py'77 else:78 raise ValueError('Unknown dataset: {}'.format(args.dataset))79 with read_write(cfg):80 cfg.evaluate.seg.cfg = seg_cfg81 cfg.evaluate.seg.opts = ['test_cfg.mode=whole']82 83 dataset_cfg = mmcv.Config()84 dataset_cfg.CLASSES = list(dataset_class.CLASSES)85 dataset_cfg.PALETTE = dataset_class.PALETTE.copy()86 87 if len(additional_classes) > 0:88 additional_classes = additional_classes.split(',')89 additional_classes = list(90 set(additional_classes) - set(dataset_cfg.CLASSES))91 dataset_cfg.CLASSES.extend(additional_classes)92 dataset_cfg.PALETTE.extend(GROUP_PALETTE[np.random.choice(93 list(range(len(GROUP_PALETTE))), len(additional_classes))])94 seg_model = build_seg_inference(model, dataset_cfg, text_transform,95 cfg.evaluate.seg)96 97 device = next(seg_model.parameters()).device98 # prepare data99 data = dict(img=input_img)100 data = test_pipeline(data)101 data = collate([data], samples_per_gpu=1)102 if next(seg_model.parameters()).is_cuda:103 # scatter to specified GPU104 data = scatter(data, [device])[0]105 else:106 data['img_metas'] = [i.data[0] for i in data['img_metas']]107 with torch.no_grad():108 result = seg_model(return_loss=False, rescale=False, **data)109 110 img_tensor = data['img'][0]111 img_metas = data['img_metas'][0]112 imgs = tensor2imgs(img_tensor, **img_metas[0]['img_norm_cfg'])113 assert len(imgs) == len(img_metas)114 115 out_file_dict = dict()116 for img, img_meta in zip(imgs, img_metas):117 h, w, _ = img_meta['img_shape']118 img_show = img[:h, :w, :]119 120 # ori_h, ori_w = img_meta['ori_shape'][:-1]121 122 # short_side = 448123 # if ori_h > ori_w:124 # new_h, new_w = ori_h * short_side//ori_w , short_side125 # else:126 # new_w, new_h = ori_w * short_side//ori_h , short_side127 128 # img_show = mmcv.imresize(img_show, (new_w, new_h))129 130 for vis_mode in vis_modes:131 out_file = osp.join(output_dir, 'vis_imgs', vis_mode,132 f'{vis_mode}.jpg')133 seg_model.show_result(img_show, img_tensor.to(device), result,134 out_file, vis_mode)135 out_file_dict[vis_mode] = out_file136 137 return [out_file_dict[mode] for mode in vis_modes]138 139 140title = 'GroupViT'141 142description = """143Gradio Demo for GroupViT: Semantic Segmentation Emerges from Text Supervision. \n144You may click on of the examples or upload your own image. \n145GroupViT could perform open vocabulary segmentation, you may input more classes (seperate by comma).146"""147 148article = """149<p style='text-align: center'>150<a href='https://arxiv.org/abs/2202.11094' target='_blank'>151GroupViT: Semantic Segmentation Emerges from Text Supervision152</a>153|154<a href='https://github.com/NVlabs/GroupViT' target='_blank'>Github Repo</a></p>155"""156 157gr.Interface(158 inference,159 inputs=[160 gr.inputs.Dropdown(dataset_options, type='value', label='Category list'),161 gr.inputs.Textbox(162 lines=1, placeholder=None, default='', label='More classes'),163 gr.inputs.Image(type='filepath')164 ],165 outputs=[gr.outputs.Image(label=label) for label in output_labels],166 title=title,167 description=description,168 article=article,169 examples=examples).launch(enable_queue=True)170 