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