Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
hungarian_tracker.py172 linesDownload Raw Back to tracking
1#!/usr/bin/env python32# Copyright 2004-present Facebook. All Rights Reserved.3import copy4import numpy as np5from typing import Dict6import torch7from scipy.optimize import linear_sum_assignment8 9from detectron2.config import configurable10from detectron2.structures import Boxes, Instances11 12from ..config.config import CfgNode as CfgNode_13from .base_tracker import BaseTracker14 15 16class BaseHungarianTracker(BaseTracker):17    """18    A base class for all Hungarian trackers19    """20 21    @configurable22    def __init__(23        self,24        video_height: int,25        video_width: int,26        max_num_instances: int = 200,27        max_lost_frame_count: int = 0,28        min_box_rel_dim: float = 0.02,29        min_instance_period: int = 1,30        **kwargs31    ):32        """33        Args:34        video_height: height the video frame35        video_width: width of the video frame36        max_num_instances: maximum number of id allowed to be tracked37        max_lost_frame_count: maximum number of frame an id can lost tracking38                              exceed this number, an id is considered as lost39                              forever40        min_box_rel_dim: a percentage, smaller than this dimension, a bbox is41                         removed from tracking42        min_instance_period: an instance will be shown after this number of period43                             since its first showing up in the video44        """45        super().__init__(**kwargs)46        self._video_height = video_height47        self._video_width = video_width48        self._max_num_instances = max_num_instances49        self._max_lost_frame_count = max_lost_frame_count50        self._min_box_rel_dim = min_box_rel_dim51        self._min_instance_period = min_instance_period52 53    @classmethod54    def from_config(cls, cfg: CfgNode_) -> Dict:55        raise NotImplementedError("Calling HungarianTracker::from_config")56 57    def build_cost_matrix(self, instances: Instances, prev_instances: Instances) -> np.ndarray:58        raise NotImplementedError("Calling HungarianTracker::build_matrix")59 60    def update(self, instances: Instances) -> Instances:61        if instances.has("pred_keypoints"):62            raise NotImplementedError("Need to add support for keypoints")63        instances = self._initialize_extra_fields(instances)64        if self._prev_instances is not None:65            self._untracked_prev_idx = set(range(len(self._prev_instances)))66            cost_matrix = self.build_cost_matrix(instances, self._prev_instances)67            matched_idx, matched_prev_idx = linear_sum_assignment(cost_matrix)68            instances = self._process_matched_idx(instances, matched_idx, matched_prev_idx)69            instances = self._process_unmatched_idx(instances, matched_idx)70            instances = self._process_unmatched_prev_idx(instances, matched_prev_idx)71        self._prev_instances = copy.deepcopy(instances)72        return instances73 74    def _initialize_extra_fields(self, instances: Instances) -> Instances:75        """76        If input instances don't have ID, ID_period, lost_frame_count fields,77        this method is used to initialize these fields.78 79        Args:80            instances: D2 Instances, for predictions of the current frame81        Return:82            D2 Instances with extra fields added83        """84        if not instances.has("ID"):85            instances.set("ID", [None] * len(instances))86        if not instances.has("ID_period"):87            instances.set("ID_period", [None] * len(instances))88        if not instances.has("lost_frame_count"):89            instances.set("lost_frame_count", [None] * len(instances))90        if self._prev_instances is None:91            instances.ID = list(range(len(instances)))92            self._id_count += len(instances)93            instances.ID_period = [1] * len(instances)94            instances.lost_frame_count = [0] * len(instances)95        return instances96 97    def _process_matched_idx(98        self, instances: Instances, matched_idx: np.ndarray, matched_prev_idx: np.ndarray99    ) -> Instances:100        assert matched_idx.size == matched_prev_idx.size101        for i in range(matched_idx.size):102            instances.ID[matched_idx[i]] = self._prev_instances.ID[matched_prev_idx[i]]103            instances.ID_period[matched_idx[i]] = (104                self._prev_instances.ID_period[matched_prev_idx[i]] + 1105            )106            instances.lost_frame_count[matched_idx[i]] = 0107        return instances108 109    def _process_unmatched_idx(self, instances: Instances, matched_idx: np.ndarray) -> Instances:110        untracked_idx = set(range(len(instances))).difference(set(matched_idx))111        for idx in untracked_idx:112            instances.ID[idx] = self._id_count113            self._id_count += 1114            instances.ID_period[idx] = 1115            instances.lost_frame_count[idx] = 0116        return instances117 118    def _process_unmatched_prev_idx(119        self, instances: Instances, matched_prev_idx: np.ndarray120    ) -> Instances:121        untracked_instances = Instances(122            image_size=instances.image_size,123            pred_boxes=[],124            pred_masks=[],125            pred_classes=[],126            scores=[],127            ID=[],128            ID_period=[],129            lost_frame_count=[],130        )131        prev_bboxes = list(self._prev_instances.pred_boxes)132        prev_classes = list(self._prev_instances.pred_classes)133        prev_scores = list(self._prev_instances.scores)134        prev_ID_period = self._prev_instances.ID_period135        if instances.has("pred_masks"):136            prev_masks = list(self._prev_instances.pred_masks)137        untracked_prev_idx = set(range(len(self._prev_instances))).difference(set(matched_prev_idx))138        for idx in untracked_prev_idx:139            x_left, y_top, x_right, y_bot = prev_bboxes[idx]140            if (141                (1.0 * (x_right - x_left) / self._video_width < self._min_box_rel_dim)142                or (1.0 * (y_bot - y_top) / self._video_height < self._min_box_rel_dim)143                or self._prev_instances.lost_frame_count[idx] >= self._max_lost_frame_count144                or prev_ID_period[idx] <= self._min_instance_period145            ):146                continue147            untracked_instances.pred_boxes.append(list(prev_bboxes[idx].numpy()))148            untracked_instances.pred_classes.append(int(prev_classes[idx]))149            untracked_instances.scores.append(float(prev_scores[idx]))150            untracked_instances.ID.append(self._prev_instances.ID[idx])151            untracked_instances.ID_period.append(self._prev_instances.ID_period[idx])152            untracked_instances.lost_frame_count.append(153                self._prev_instances.lost_frame_count[idx] + 1154            )155            if instances.has("pred_masks"):156                untracked_instances.pred_masks.append(prev_masks[idx].numpy().astype(np.uint8))157 158        untracked_instances.pred_boxes = Boxes(torch.FloatTensor(untracked_instances.pred_boxes))159        untracked_instances.pred_classes = torch.IntTensor(untracked_instances.pred_classes)160        untracked_instances.scores = torch.FloatTensor(untracked_instances.scores)161        if instances.has("pred_masks"):162            untracked_instances.pred_masks = torch.IntTensor(untracked_instances.pred_masks)163        else:164            untracked_instances.remove("pred_masks")165 166        return Instances.cat(167            [168                instances,169                untracked_instances,170            ]171        )172