Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
bbox_iou_tracker.py277 linesDownload Raw Back to tracking
1#!/usr/bin/env python32# Copyright 2004-present Facebook. All Rights Reserved.3import copy4import numpy as np5from typing import List6import torch7 8from detectron2.config import configurable9from detectron2.structures import Boxes, Instances10from detectron2.structures.boxes import pairwise_iou11 12from ..config.config import CfgNode as CfgNode_13from .base_tracker import TRACKER_HEADS_REGISTRY, BaseTracker14 15 16@TRACKER_HEADS_REGISTRY.register()17class BBoxIOUTracker(BaseTracker):18    """19    A bounding box tracker to assign ID based on IoU between current and previous instances20    """21 22    @configurable23    def __init__(24        self,25        *,26        video_height: int,27        video_width: int,28        max_num_instances: int = 200,29        max_lost_frame_count: int = 0,30        min_box_rel_dim: float = 0.02,31        min_instance_period: int = 1,32        track_iou_threshold: float = 0.5,33        **kwargs,34    ):35        """36        Args:37        video_height: height the video frame38        video_width: width of the video frame39        max_num_instances: maximum number of id allowed to be tracked40        max_lost_frame_count: maximum number of frame an id can lost tracking41                              exceed this number, an id is considered as lost42                              forever43        min_box_rel_dim: a percentage, smaller than this dimension, a bbox is44                         removed from tracking45        min_instance_period: an instance will be shown after this number of period46                             since its first showing up in the video47        track_iou_threshold: iou threshold, below this number a bbox pair is removed48                             from tracking49        """50        super().__init__(**kwargs)51        self._video_height = video_height52        self._video_width = video_width53        self._max_num_instances = max_num_instances54        self._max_lost_frame_count = max_lost_frame_count55        self._min_box_rel_dim = min_box_rel_dim56        self._min_instance_period = min_instance_period57        self._track_iou_threshold = track_iou_threshold58 59    @classmethod60    def from_config(cls, cfg: CfgNode_):61        """62        Old style initialization using CfgNode63 64        Args:65            cfg: D2 CfgNode, config file66        Return:67            dictionary storing arguments for __init__ method68        """69        assert "VIDEO_HEIGHT" in cfg.TRACKER_HEADS70        assert "VIDEO_WIDTH" in cfg.TRACKER_HEADS71        video_height = cfg.TRACKER_HEADS.get("VIDEO_HEIGHT")72        video_width = cfg.TRACKER_HEADS.get("VIDEO_WIDTH")73        max_num_instances = cfg.TRACKER_HEADS.get("MAX_NUM_INSTANCES", 200)74        max_lost_frame_count = cfg.TRACKER_HEADS.get("MAX_LOST_FRAME_COUNT", 0)75        min_box_rel_dim = cfg.TRACKER_HEADS.get("MIN_BOX_REL_DIM", 0.02)76        min_instance_period = cfg.TRACKER_HEADS.get("MIN_INSTANCE_PERIOD", 1)77        track_iou_threshold = cfg.TRACKER_HEADS.get("TRACK_IOU_THRESHOLD", 0.5)78        return {79            "_target_": "detectron2.tracking.bbox_iou_tracker.BBoxIOUTracker",80            "video_height": video_height,81            "video_width": video_width,82            "max_num_instances": max_num_instances,83            "max_lost_frame_count": max_lost_frame_count,84            "min_box_rel_dim": min_box_rel_dim,85            "min_instance_period": min_instance_period,86            "track_iou_threshold": track_iou_threshold,87        }88 89    def update(self, instances: Instances) -> Instances:90        """91        See BaseTracker description92        """93        instances = self._initialize_extra_fields(instances)94        if self._prev_instances is not None:95            # calculate IoU of all bbox pairs96            iou_all = pairwise_iou(97                boxes1=instances.pred_boxes,98                boxes2=self._prev_instances.pred_boxes,99            )100            # sort IoU in descending order101            bbox_pairs = self._create_prediction_pairs(instances, iou_all)102            # assign previous ID to current bbox if IoU > track_iou_threshold103            self._reset_fields()104            for bbox_pair in bbox_pairs:105                idx = bbox_pair["idx"]106                prev_id = bbox_pair["prev_id"]107                if (108                    idx in self._matched_idx109                    or prev_id in self._matched_ID110                    or bbox_pair["IoU"] < self._track_iou_threshold111                ):112                    continue113                instances.ID[idx] = prev_id114                instances.ID_period[idx] = bbox_pair["prev_period"] + 1115                instances.lost_frame_count[idx] = 0116                self._matched_idx.add(idx)117                self._matched_ID.add(prev_id)118                self._untracked_prev_idx.remove(bbox_pair["prev_idx"])119            instances = self._assign_new_id(instances)120            instances = self._merge_untracked_instances(instances)121        self._prev_instances = copy.deepcopy(instances)122        return instances123 124    def _create_prediction_pairs(self, instances: Instances, iou_all: np.ndarray) -> List:125        """126        For all instances in previous and current frames, create pairs. For each127        pair, store index of the instance in current frame predcitions, index in128        previous predictions, ID in previous predictions, IoU of the bboxes in this129        pair, period in previous predictions.130 131        Args:132            instances: D2 Instances, for predictions of the current frame133            iou_all: IoU for all bboxes pairs134        Return:135            A list of IoU for all pairs136        """137        bbox_pairs = []138        for i in range(len(instances)):139            for j in range(len(self._prev_instances)):140                bbox_pairs.append(141                    {142                        "idx": i,143                        "prev_idx": j,144                        "prev_id": self._prev_instances.ID[j],145                        "IoU": iou_all[i, j],146                        "prev_period": self._prev_instances.ID_period[j],147                    }148                )149        return bbox_pairs150 151    def _initialize_extra_fields(self, instances: Instances) -> Instances:152        """153        If input instances don't have ID, ID_period, lost_frame_count fields,154        this method is used to initialize these fields.155 156        Args:157            instances: D2 Instances, for predictions of the current frame158        Return:159            D2 Instances with extra fields added160        """161        if not instances.has("ID"):162            instances.set("ID", [None] * len(instances))163        if not instances.has("ID_period"):164            instances.set("ID_period", [None] * len(instances))165        if not instances.has("lost_frame_count"):166            instances.set("lost_frame_count", [None] * len(instances))167        if self._prev_instances is None:168            instances.ID = list(range(len(instances)))169            self._id_count += len(instances)170            instances.ID_period = [1] * len(instances)171            instances.lost_frame_count = [0] * len(instances)172        return instances173 174    def _reset_fields(self):175        """176        Before each uodate call, reset fields first177        """178        self._matched_idx = set()179        self._matched_ID = set()180        self._untracked_prev_idx = set(range(len(self._prev_instances)))181 182    def _assign_new_id(self, instances: Instances) -> Instances:183        """184        For each untracked instance, assign a new id185 186        Args:187            instances: D2 Instances, for predictions of the current frame188        Return:189            D2 Instances with new ID assigned190        """191        untracked_idx = set(range(len(instances))).difference(self._matched_idx)192        for idx in untracked_idx:193            instances.ID[idx] = self._id_count194            self._id_count += 1195            instances.ID_period[idx] = 1196            instances.lost_frame_count[idx] = 0197        return instances198 199    def _merge_untracked_instances(self, instances: Instances) -> Instances:200        """201        For untracked previous instances, under certain condition, still keep them202        in tracking and merge with the current instances.203 204        Args:205            instances: D2 Instances, for predictions of the current frame206        Return:207            D2 Instances merging current instances and instances from previous208            frame decided to keep tracking209        """210        untracked_instances = Instances(211            image_size=instances.image_size,212            pred_boxes=[],213            pred_classes=[],214            scores=[],215            ID=[],216            ID_period=[],217            lost_frame_count=[],218        )219        prev_bboxes = list(self._prev_instances.pred_boxes)220        prev_classes = list(self._prev_instances.pred_classes)221        prev_scores = list(self._prev_instances.scores)222        prev_ID_period = self._prev_instances.ID_period223        if instances.has("pred_masks"):224            untracked_instances.set("pred_masks", [])225            prev_masks = list(self._prev_instances.pred_masks)226        if instances.has("pred_keypoints"):227            untracked_instances.set("pred_keypoints", [])228            prev_keypoints = list(self._prev_instances.pred_keypoints)229        if instances.has("pred_keypoint_heatmaps"):230            untracked_instances.set("pred_keypoint_heatmaps", [])231            prev_keypoint_heatmaps = list(self._prev_instances.pred_keypoint_heatmaps)232        for idx in self._untracked_prev_idx:233            x_left, y_top, x_right, y_bot = prev_bboxes[idx]234            if (235                (1.0 * (x_right - x_left) / self._video_width < self._min_box_rel_dim)236                or (1.0 * (y_bot - y_top) / self._video_height < self._min_box_rel_dim)237                or self._prev_instances.lost_frame_count[idx] >= self._max_lost_frame_count238                or prev_ID_period[idx] <= self._min_instance_period239            ):240                continue241            untracked_instances.pred_boxes.append(list(prev_bboxes[idx].numpy()))242            untracked_instances.pred_classes.append(int(prev_classes[idx]))243            untracked_instances.scores.append(float(prev_scores[idx]))244            untracked_instances.ID.append(self._prev_instances.ID[idx])245            untracked_instances.ID_period.append(self._prev_instances.ID_period[idx])246            untracked_instances.lost_frame_count.append(247                self._prev_instances.lost_frame_count[idx] + 1248            )249            if instances.has("pred_masks"):250                untracked_instances.pred_masks.append(prev_masks[idx].numpy().astype(np.uint8))251            if instances.has("pred_keypoints"):252                untracked_instances.pred_keypoints.append(253                    prev_keypoints[idx].numpy().astype(np.uint8)254                )255            if instances.has("pred_keypoint_heatmaps"):256                untracked_instances.pred_keypoint_heatmaps.append(257                    prev_keypoint_heatmaps[idx].numpy().astype(np.float32)258                )259        untracked_instances.pred_boxes = Boxes(torch.FloatTensor(untracked_instances.pred_boxes))260        untracked_instances.pred_classes = torch.IntTensor(untracked_instances.pred_classes)261        untracked_instances.scores = torch.FloatTensor(untracked_instances.scores)262        if instances.has("pred_masks"):263            untracked_instances.pred_masks = torch.IntTensor(untracked_instances.pred_masks)264        if instances.has("pred_keypoints"):265            untracked_instances.pred_keypoints = torch.IntTensor(untracked_instances.pred_keypoints)266        if instances.has("pred_keypoint_heatmaps"):267            untracked_instances.pred_keypoint_heatmaps = torch.FloatTensor(268                untracked_instances.pred_keypoint_heatmaps269            )270 271        return Instances.cat(272            [273                instances,274                untracked_instances,275            ]276        )277