Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
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 