Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
video_visualizer.py288 linesDownload Raw Back to utils
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