Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2import numpy as np3from typing import List4import pycocotools.mask as mask_util5 6from detectron2.structures import Instances7from detectron2.utils.visualizer import (8 ColorMode,9 Visualizer,10 _create_text_labels,11 _PanopticPrediction,12)13 14from .colormap import random_color, random_colors15 16 17class _DetectedInstance:18 """19 Used to store data about detected objects in video frame,20 in order to transfer color to objects in the future frames.21 22 Attributes:23 label (int):24 bbox (tuple[float]):25 mask_rle (dict):26 color (tuple[float]): RGB colors in range (0, 1)27 ttl (int): time-to-live for the instance. For example, if ttl=2,28 the instance color can be transferred to objects in the next two frames.29 """30 31 __slots__ = ["label", "bbox", "mask_rle", "color", "ttl"]32 33 def __init__(self, label, bbox, mask_rle, color, ttl):34 self.label = label35 self.bbox = bbox36 self.mask_rle = mask_rle37 self.color = color38 self.ttl = ttl39 40 41class VideoVisualizer:42 def __init__(self, metadata, instance_mode=ColorMode.IMAGE):43 """44 Args:45 metadata (MetadataCatalog): image metadata.46 """47 self.metadata = metadata48 self._old_instances = []49 assert instance_mode in [50 ColorMode.IMAGE,51 ColorMode.IMAGE_BW,52 ], "Other mode not supported yet."53 self._instance_mode = instance_mode54 self._max_num_instances = self.metadata.get("max_num_instances", 74)55 self._assigned_colors = {}56 self._color_pool = random_colors(self._max_num_instances, rgb=True, maximum=1)57 self._color_idx_set = set(range(len(self._color_pool)))58 59 def draw_instance_predictions(self, frame, predictions):60 """61 Draw instance-level prediction results on an image.62 63 Args:64 frame (ndarray): an RGB image of shape (H, W, C), in the range [0, 255].65 predictions (Instances): the output of an instance detection/segmentation66 model. Following fields will be used to draw:67 "pred_boxes", "pred_classes", "scores", "pred_masks" (or "pred_masks_rle").68 69 Returns:70 output (VisImage): image object with visualizations.71 """72 frame_visualizer = Visualizer(frame, self.metadata)73 num_instances = len(predictions)74 if num_instances == 0:75 return frame_visualizer.output76 77 boxes = predictions.pred_boxes.tensor.numpy() if predictions.has("pred_boxes") else None78 scores = predictions.scores if predictions.has("scores") else None79 classes = predictions.pred_classes.numpy() if predictions.has("pred_classes") else None80 keypoints = predictions.pred_keypoints if predictions.has("pred_keypoints") else None81 colors = predictions.COLOR if predictions.has("COLOR") else [None] * len(predictions)82 periods = predictions.ID_period if predictions.has("ID_period") else None83 period_threshold = self.metadata.get("period_threshold", 0)84 visibilities = (85 [True] * len(predictions)86 if periods is None87 else [x > period_threshold for x in periods]88 )89 90 if predictions.has("pred_masks"):91 masks = predictions.pred_masks92 # mask IOU is not yet enabled93 # masks_rles = mask_util.encode(np.asarray(masks.permute(1, 2, 0), order="F"))94 # assert len(masks_rles) == num_instances95 else:96 masks = None97 98 if not predictions.has("COLOR"):99 if predictions.has("ID"):100 colors = self._assign_colors_by_id(predictions)101 else:102 # ToDo: clean old assign color method and use a default tracker to assign id103 detected = [104 _DetectedInstance(classes[i], boxes[i], mask_rle=None, color=colors[i], ttl=8)105 for i in range(num_instances)106 ]107 colors = self._assign_colors(detected)108 109 labels = _create_text_labels(classes, scores, self.metadata.get("thing_classes", None))110 111 if self._instance_mode == ColorMode.IMAGE_BW:112 # any() returns uint8 tensor113 frame_visualizer.output.reset_image(114 frame_visualizer._create_grayscale_image(115 (masks.any(dim=0) > 0).numpy() if masks is not None else None116 )117 )118 alpha = 0.3119 else:120 alpha = 0.5121 122 labels = (123 None124 if labels is None125 else [y[0] for y in filter(lambda x: x[1], zip(labels, visibilities))]126 ) # noqa127 assigned_colors = (128 None129 if colors is None130 else [y[0] for y in filter(lambda x: x[1], zip(colors, visibilities))]131 ) # noqa132 frame_visualizer.overlay_instances(133 boxes=None if masks is not None else boxes[visibilities], # boxes are a bit distracting134 masks=None if masks is None else masks[visibilities],135 labels=labels,136 keypoints=None if keypoints is None else keypoints[visibilities],137 assigned_colors=assigned_colors,138 alpha=alpha,139 )140 141 return frame_visualizer.output142 143 def draw_sem_seg(self, frame, sem_seg, area_threshold=None):144 """145 Args:146 sem_seg (ndarray or Tensor): semantic segmentation of shape (H, W),147 each value is the integer label.148 area_threshold (Optional[int]): only draw segmentations larger than the threshold149 """150 # don't need to do anything special151 frame_visualizer = Visualizer(frame, self.metadata)152 frame_visualizer.draw_sem_seg(sem_seg, area_threshold=None)153 return frame_visualizer.output154 155 def draw_panoptic_seg_predictions(156 self, frame, panoptic_seg, segments_info, area_threshold=None, alpha=0.5157 ):158 frame_visualizer = Visualizer(frame, self.metadata)159 pred = _PanopticPrediction(panoptic_seg, segments_info, self.metadata)160 161 if self._instance_mode == ColorMode.IMAGE_BW:162 frame_visualizer.output.reset_image(163 frame_visualizer._create_grayscale_image(pred.non_empty_mask())164 )165 166 # draw mask for all semantic segments first i.e. "stuff"167 for mask, sinfo in pred.semantic_masks():168 category_idx = sinfo["category_id"]169 try:170 mask_color = [x / 255 for x in self.metadata.stuff_colors[category_idx]]171 except AttributeError:172 mask_color = None173 174 frame_visualizer.draw_binary_mask(175 mask,176 color=mask_color,177 text=self.metadata.stuff_classes[category_idx],178 alpha=alpha,179 area_threshold=area_threshold,180 )181 182 all_instances = list(pred.instance_masks())183 if len(all_instances) == 0:184 return frame_visualizer.output185 # draw mask for all instances second186 masks, sinfo = list(zip(*all_instances))187 num_instances = len(masks)188 masks_rles = mask_util.encode(189 np.asarray(np.asarray(masks).transpose(1, 2, 0), dtype=np.uint8, order="F")190 )191 assert len(masks_rles) == num_instances192 193 category_ids = [x["category_id"] for x in sinfo]194 detected = [195 _DetectedInstance(category_ids[i], bbox=None, mask_rle=masks_rles[i], color=None, ttl=8)196 for i in range(num_instances)197 ]198 colors = self._assign_colors(detected)199 labels = [self.metadata.thing_classes[k] for k in category_ids]200 201 frame_visualizer.overlay_instances(202 boxes=None,203 masks=masks,204 labels=labels,205 keypoints=None,206 assigned_colors=colors,207 alpha=alpha,208 )209 return frame_visualizer.output210 211 def _assign_colors(self, instances):212 """213 Naive tracking heuristics to assign same color to the same instance,214 will update the internal state of tracked instances.215 216 Returns:217 list[tuple[float]]: list of colors.218 """219 220 # Compute iou with either boxes or masks:221 is_crowd = np.zeros((len(instances),), dtype=bool)222 if instances[0].bbox is None:223 assert instances[0].mask_rle is not None224 # use mask iou only when box iou is None225 # because box seems good enough226 rles_old = [x.mask_rle for x in self._old_instances]227 rles_new = [x.mask_rle for x in instances]228 ious = mask_util.iou(rles_old, rles_new, is_crowd)229 threshold = 0.5230 else:231 boxes_old = [x.bbox for x in self._old_instances]232 boxes_new = [x.bbox for x in instances]233 ious = mask_util.iou(boxes_old, boxes_new, is_crowd)234 threshold = 0.6235 if len(ious) == 0:236 ious = np.zeros((len(self._old_instances), len(instances)), dtype="float32")237 238 # Only allow matching instances of the same label:239 for old_idx, old in enumerate(self._old_instances):240 for new_idx, new in enumerate(instances):241 if old.label != new.label:242 ious[old_idx, new_idx] = 0243 244 matched_new_per_old = np.asarray(ious).argmax(axis=1)245 max_iou_per_old = np.asarray(ious).max(axis=1)246 247 # Try to find match for each old instance:248 extra_instances = []249 for idx, inst in enumerate(self._old_instances):250 if max_iou_per_old[idx] > threshold:251 newidx = matched_new_per_old[idx]252 if instances[newidx].color is None:253 instances[newidx].color = inst.color254 continue255 # If an old instance does not match any new instances,256 # keep it for the next frame in case it is just missed by the detector257 inst.ttl -= 1258 if inst.ttl > 0:259 extra_instances.append(inst)260 261 # Assign random color to newly-detected instances:262 for inst in instances:263 if inst.color is None:264 inst.color = random_color(rgb=True, maximum=1)265 self._old_instances = instances[:] + extra_instances266 return [d.color for d in instances]267 268 def _assign_colors_by_id(self, instances: Instances) -> List:269 colors = []270 untracked_ids = set(self._assigned_colors.keys())271 for id in instances.ID:272 if id in self._assigned_colors:273 colors.append(self._color_pool[self._assigned_colors[id]])274 untracked_ids.remove(id)275 else:276 assert (277 len(self._color_idx_set) >= 1278 ), f"Number of id exceeded maximum, \279 max = {self._max_num_instances}"280 idx = self._color_idx_set.pop()281 color = self._color_pool[idx]282 self._assigned_colors[id] = idx283 colors.append(color)284 for id in untracked_ids:285 self._color_idx_set.add(self._assigned_colors[id])286 del self._assigned_colors[id]287 return colors288 