xdecoder/Instruct-X-Decoder
163
1# --------------------------------------------------------2# X-Decoder -- Generalized Decoding for Pixel, Image, and Language3# Copyright (c) 2022 Microsoft4# Licensed under The MIT License [see LICENSE for details]5# Written by Xueyan Zou (xueyan@cs.wisc.edu)6# --------------------------------------------------------7 8import torch9import numpy as np10from PIL import Image11from torchvision import transforms12from utils.visualizer import Visualizer13from detectron2.utils.colormap import random_color14from detectron2.data import MetadataCatalog15 16 17t = []18t.append(transforms.Resize(512, interpolation=Image.BICUBIC))19transform = transforms.Compose(t)20metadata = MetadataCatalog.get('ade20k_panoptic_train')21 22def referring_segmentation(model, image, texts, inpainting_text, *args, **kwargs):23 model.model.metadata = metadata24 texts = texts.strip()25 texts = [[text.strip() if text.endswith('.') else (text + '.')] for text in texts.split(',')]26 image_ori = transform(image)27 28 with torch.no_grad():29 width = image_ori.size[0]30 height = image_ori.size[1]31 image = np.asarray(image_ori)32 image_ori_np = np.asarray(image_ori)33 images = torch.from_numpy(image.copy()).permute(2,0,1).cuda()34 35 batch_inputs = [{'image': images, 'height': height, 'width': width, 'groundings': {'texts': texts}}] 36 outputs = model.model.evaluate_grounding(batch_inputs, None)37 visual = Visualizer(image_ori_np, metadata=metadata)38 39 grd_mask = (outputs[0]['grounding_mask'] > 0).float().cpu().numpy()40 for idx, mask in enumerate(grd_mask):41 color = random_color(rgb=True, maximum=1).astype(np.int32).tolist()42 demo = visual.draw_binary_mask(mask, color=color, text=texts[idx])43 res = demo.get_image()44 45 torch.cuda.empty_cache()46 return Image.fromarray(res), '', None