Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
instances.py195 linesDownload Raw Back to structures
1# Copyright (c) Facebook, Inc. and its affiliates.2import itertools3import warnings4from typing import Any, Dict, List, Tuple, Union5import torch6 7 8class Instances:9    """10    This class represents a list of instances in an image.11    It stores the attributes of instances (e.g., boxes, masks, labels, scores) as "fields".12    All fields must have the same ``__len__`` which is the number of instances.13 14    All other (non-field) attributes of this class are considered private:15    they must start with '_' and are not modifiable by a user.16 17    Some basic usage:18 19    1. Set/get/check a field:20 21       .. code-block:: python22 23          instances.gt_boxes = Boxes(...)24          print(instances.pred_masks)  # a tensor of shape (N, H, W)25          print('gt_masks' in instances)26 27    2. ``len(instances)`` returns the number of instances28    3. Indexing: ``instances[indices]`` will apply the indexing on all the fields29       and returns a new :class:`Instances`.30       Typically, ``indices`` is a integer vector of indices,31       or a binary mask of length ``num_instances``32 33       .. code-block:: python34 35          category_3_detections = instances[instances.pred_classes == 3]36          confident_detections = instances[instances.scores > 0.9]37    """38 39    def __init__(self, image_size: Tuple[int, int], **kwargs: Any):40        """41        Args:42            image_size (height, width): the spatial size of the image.43            kwargs: fields to add to this `Instances`.44        """45        self._image_size = image_size46        self._fields: Dict[str, Any] = {}47        for k, v in kwargs.items():48            self.set(k, v)49 50    @property51    def image_size(self) -> Tuple[int, int]:52        """53        Returns:54            tuple: height, width55        """56        return self._image_size57 58    def __setattr__(self, name: str, val: Any) -> None:59        if name.startswith("_"):60            super().__setattr__(name, val)61        else:62            self.set(name, val)63 64    def __getattr__(self, name: str) -> Any:65        if name == "_fields" or name not in self._fields:66            raise AttributeError("Cannot find field '{}' in the given Instances!".format(name))67        return self._fields[name]68 69    def set(self, name: str, value: Any) -> None:70        """71        Set the field named `name` to `value`.72        The length of `value` must be the number of instances,73        and must agree with other existing fields in this object.74        """75        with warnings.catch_warnings(record=True):76            data_len = len(value)77        if len(self._fields):78            assert (79                len(self) == data_len80            ), "Adding a field of length {} to a Instances of length {}".format(data_len, len(self))81        self._fields[name] = value82 83    def has(self, name: str) -> bool:84        """85        Returns:86            bool: whether the field called `name` exists.87        """88        return name in self._fields89 90    def remove(self, name: str) -> None:91        """92        Remove the field called `name`.93        """94        del self._fields[name]95 96    def get(self, name: str) -> Any:97        """98        Returns the field called `name`.99        """100        return self._fields[name]101 102    def get_fields(self) -> Dict[str, Any]:103        """104        Returns:105            dict: a dict which maps names (str) to data of the fields106 107        Modifying the returned dict will modify this instance.108        """109        return self._fields110 111    # Tensor-like methods112    def to(self, *args: Any, **kwargs: Any) -> "Instances":113        """114        Returns:115            Instances: all fields are called with a `to(device)`, if the field has this method.116        """117        ret = Instances(self._image_size)118        for k, v in self._fields.items():119            if hasattr(v, "to"):120                v = v.to(*args, **kwargs)121            ret.set(k, v)122        return ret123 124    def __getitem__(self, item: Union[int, slice, torch.BoolTensor]) -> "Instances":125        """126        Args:127            item: an index-like object and will be used to index all the fields.128 129        Returns:130            If `item` is a string, return the data in the corresponding field.131            Otherwise, returns an `Instances` where all fields are indexed by `item`.132        """133        if type(item) == int:134            if item >= len(self) or item < -len(self):135                raise IndexError("Instances index out of range!")136            else:137                item = slice(item, None, len(self))138 139        ret = Instances(self._image_size)140        for k, v in self._fields.items():141            ret.set(k, v[item])142        return ret143 144    def __len__(self) -> int:145        for v in self._fields.values():146            # use __len__ because len() has to be int and is not friendly to tracing147            return v.__len__()148        raise NotImplementedError("Empty Instances does not support __len__!")149 150    def __iter__(self):151        raise NotImplementedError("`Instances` object is not iterable!")152 153    @staticmethod154    def cat(instance_lists: List["Instances"]) -> "Instances":155        """156        Args:157            instance_lists (list[Instances])158 159        Returns:160            Instances161        """162        assert all(isinstance(i, Instances) for i in instance_lists)163        assert len(instance_lists) > 0164        if len(instance_lists) == 1:165            return instance_lists[0]166 167        image_size = instance_lists[0].image_size168        if not isinstance(image_size, torch.Tensor):  # could be a tensor in tracing169            for i in instance_lists[1:]:170                assert i.image_size == image_size171        ret = Instances(image_size)172        for k in instance_lists[0]._fields.keys():173            values = [i.get(k) for i in instance_lists]174            v0 = values[0]175            if isinstance(v0, torch.Tensor):176                values = torch.cat(values, dim=0)177            elif isinstance(v0, list):178                values = list(itertools.chain(*values))179            elif hasattr(type(v0), "cat"):180                values = type(v0).cat(values)181            else:182                raise ValueError("Unsupported type {} for concatenation".format(type(v0)))183            ret.set(k, values)184        return ret185 186    def __str__(self) -> str:187        s = self.__class__.__name__ + "("188        s += "num_instances={}, ".format(len(self))189        s += "image_height={}, ".format(self._image_size[0])190        s += "image_width={}, ".format(self._image_size[1])191        s += "fields=[{}])".format(", ".join((f"{k}: {v}" for k, v in self._fields.items())))192        return s193 194    __repr__ = __str__195