Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
base_tracker.py65 linesDownload Raw Back to tracking
1#!/usr/bin/env python32# Copyright 2004-present Facebook. All Rights Reserved.3from detectron2.config import configurable4from detectron2.utils.registry import Registry5 6from ..config.config import CfgNode as CfgNode_7from ..structures import Instances8 9TRACKER_HEADS_REGISTRY = Registry("TRACKER_HEADS")10TRACKER_HEADS_REGISTRY.__doc__ = """11Registry for tracking classes.12"""13 14 15class BaseTracker(object):16    """17    A parent class for all trackers18    """19 20    @configurable21    def __init__(self, **kwargs):22        self._prev_instances = None  # (D2)instances for previous frame23        self._matched_idx = set()  # indices in prev_instances found matching24        self._matched_ID = set()  # idendities in prev_instances found matching25        self._untracked_prev_idx = set()  # indices in prev_instances not found matching26        self._id_count = 0  # used to assign new id27 28    @classmethod29    def from_config(cls, cfg: CfgNode_):30        raise NotImplementedError("Calling BaseTracker::from_config")31 32    def update(self, predictions: Instances) -> Instances:33        """34        Args:35            predictions: D2 Instances for predictions of the current frame36        Return:37            D2 Instances for predictions of the current frame with ID assigned38 39        _prev_instances and instances will have the following fields:40          .pred_boxes               (shape=[N, 4])41          .scores                   (shape=[N,])42          .pred_classes             (shape=[N,])43          .pred_keypoints           (shape=[N, M, 3], Optional)44          .pred_masks               (shape=List[2D_MASK], Optional)   2D_MASK: shape=[H, W]45          .ID                       (shape=[N,])46 47        N: # of detected bboxes48        H and W: height and width of 2D mask49        """50        raise NotImplementedError("Calling BaseTracker::update")51 52 53def build_tracker_head(cfg: CfgNode_) -> BaseTracker:54    """55    Build a tracker head from `cfg.TRACKER_HEADS.TRACKER_NAME`.56 57    Args:58        cfg: D2 CfgNode, config file with tracker information59    Return:60        tracker object61    """62    name = cfg.TRACKER_HEADS.TRACKER_NAME63    tracker_class = TRACKER_HEADS_REGISTRY.get(name)64    return tracker_class(cfg)65