Team Ai
Apppublic

CVPR/GroupViT

sourceHugging Faceupdated 4y agoView on Hugging Face
9likes
app.py170 linesDownload Raw Back to root
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