Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2import collections3from dataclasses import dataclass4from typing import Callable, List, Optional, Tuple5import torch6from torch import nn7 8from detectron2.structures import Boxes, Instances, ROIMasks9from detectron2.utils.registry import _convert_target_to_string, locate10 11from .torchscript_patch import patch_builtin_len12 13 14@dataclass15class Schema:16 """17 A Schema defines how to flatten a possibly hierarchical object into tuple of18 primitive objects, so it can be used as inputs/outputs of PyTorch's tracing.19 20 PyTorch does not support tracing a function that produces rich output21 structures (e.g. dict, Instances, Boxes). To trace such a function, we22 flatten the rich object into tuple of tensors, and return this tuple of tensors23 instead. Meanwhile, we also need to know how to "rebuild" the original object24 from the flattened results, so we can evaluate the flattened results.25 A Schema defines how to flatten an object, and while flattening it, it records26 necessary schemas so that the object can be rebuilt using the flattened outputs.27 28 The flattened object and the schema object is returned by ``.flatten`` classmethod.29 Then the original object can be rebuilt with the ``__call__`` method of schema.30 31 A Schema is a dataclass that can be serialized easily.32 """33 34 # inspired by FetchMapper in tensorflow/python/client/session.py35 36 @classmethod37 def flatten(cls, obj):38 raise NotImplementedError39 40 def __call__(self, values):41 raise NotImplementedError42 43 @staticmethod44 def _concat(values):45 ret = ()46 sizes = []47 for v in values:48 assert isinstance(v, tuple), "Flattened results must be a tuple"49 ret = ret + v50 sizes.append(len(v))51 return ret, sizes52 53 @staticmethod54 def _split(values, sizes):55 if len(sizes):56 expected_len = sum(sizes)57 assert (58 len(values) == expected_len59 ), f"Values has length {len(values)} but expect length {expected_len}."60 ret = []61 for k in range(len(sizes)):62 begin, end = sum(sizes[:k]), sum(sizes[: k + 1])63 ret.append(values[begin:end])64 return ret65 66 67@dataclass68class ListSchema(Schema):69 schemas: List[Schema] # the schemas that define how to flatten each element in the list70 sizes: List[int] # the flattened length of each element71 72 def __call__(self, values):73 values = self._split(values, self.sizes)74 if len(values) != len(self.schemas):75 raise ValueError(76 f"Values has length {len(values)} but schemas " f"has length {len(self.schemas)}!"77 )78 values = [m(v) for m, v in zip(self.schemas, values)]79 return list(values)80 81 @classmethod82 def flatten(cls, obj):83 res = [flatten_to_tuple(k) for k in obj]84 values, sizes = cls._concat([k[0] for k in res])85 return values, cls([k[1] for k in res], sizes)86 87 88@dataclass89class TupleSchema(ListSchema):90 def __call__(self, values):91 return tuple(super().__call__(values))92 93 94@dataclass95class IdentitySchema(Schema):96 def __call__(self, values):97 return values[0]98 99 @classmethod100 def flatten(cls, obj):101 return (obj,), cls()102 103 104@dataclass105class DictSchema(ListSchema):106 keys: List[str]107 108 def __call__(self, values):109 values = super().__call__(values)110 return dict(zip(self.keys, values))111 112 @classmethod113 def flatten(cls, obj):114 for k in obj.keys():115 if not isinstance(k, str):116 raise KeyError("Only support flattening dictionaries if keys are str.")117 keys = sorted(obj.keys())118 values = [obj[k] for k in keys]119 ret, schema = ListSchema.flatten(values)120 return ret, cls(schema.schemas, schema.sizes, keys)121 122 123@dataclass124class InstancesSchema(DictSchema):125 def __call__(self, values):126 image_size, fields = values[-1], values[:-1]127 fields = super().__call__(fields)128 return Instances(image_size, **fields)129 130 @classmethod131 def flatten(cls, obj):132 ret, schema = super().flatten(obj.get_fields())133 size = obj.image_size134 if not isinstance(size, torch.Tensor):135 size = torch.tensor(size)136 return ret + (size,), schema137 138 139@dataclass140class TensorWrapSchema(Schema):141 """142 For classes that are simple wrapper of tensors, e.g.143 Boxes, RotatedBoxes, BitMasks144 """145 146 class_name: str147 148 def __call__(self, values):149 return locate(self.class_name)(values[0])150 151 @classmethod152 def flatten(cls, obj):153 return (obj.tensor,), cls(_convert_target_to_string(type(obj)))154 155 156# if more custom structures needed in the future, can allow157# passing in extra schemas for custom types158def flatten_to_tuple(obj):159 """160 Flatten an object so it can be used for PyTorch tracing.161 Also returns how to rebuild the original object from the flattened outputs.162 163 Returns:164 res (tuple): the flattened results that can be used as tracing outputs165 schema: an object with a ``__call__`` method such that ``schema(res) == obj``.166 It is a pure dataclass that can be serialized.167 """168 schemas = [169 ((str, bytes), IdentitySchema),170 (list, ListSchema),171 (tuple, TupleSchema),172 (collections.abc.Mapping, DictSchema),173 (Instances, InstancesSchema),174 ((Boxes, ROIMasks), TensorWrapSchema),175 ]176 for klass, schema in schemas:177 if isinstance(obj, klass):178 F = schema179 break180 else:181 F = IdentitySchema182 183 return F.flatten(obj)184 185 186class TracingAdapter(nn.Module):187 """188 A model may take rich input/output format (e.g. dict or custom classes),189 but `torch.jit.trace` requires tuple of tensors as input/output.190 This adapter flattens input/output format of a model so it becomes traceable.191 192 It also records the necessary schema to rebuild model's inputs/outputs from flattened193 inputs/outputs.194 195 Example:196 ::197 outputs = model(inputs) # inputs/outputs may be rich structure198 adapter = TracingAdapter(model, inputs)199 200 # can now trace the model, with adapter.flattened_inputs, or another201 # tuple of tensors with the same length and meaning202 traced = torch.jit.trace(adapter, adapter.flattened_inputs)203 204 # traced model can only produce flattened outputs (tuple of tensors)205 flattened_outputs = traced(*adapter.flattened_inputs)206 # adapter knows the schema to convert it back (new_outputs == outputs)207 new_outputs = adapter.outputs_schema(flattened_outputs)208 """209 210 flattened_inputs: Tuple[torch.Tensor] = None211 """212 Flattened version of inputs given to this class's constructor.213 """214 215 inputs_schema: Schema = None216 """217 Schema of the inputs given to this class's constructor.218 """219 220 outputs_schema: Schema = None221 """222 Schema of the output produced by calling the given model with inputs.223 """224 225 def __init__(226 self,227 model: nn.Module,228 inputs,229 inference_func: Optional[Callable] = None,230 allow_non_tensor: bool = False,231 ):232 """233 Args:234 model: an nn.Module235 inputs: An input argument or a tuple of input arguments used to call model.236 After flattening, it has to only consist of tensors.237 inference_func: a callable that takes (model, *inputs), calls the238 model with inputs, and return outputs. By default it239 is ``lambda model, *inputs: model(*inputs)``. Can be override240 if you need to call the model differently.241 allow_non_tensor: allow inputs/outputs to contain non-tensor objects.242 This option will filter out non-tensor objects to make the243 model traceable, but ``inputs_schema``/``outputs_schema`` cannot be244 used anymore because inputs/outputs cannot be rebuilt from pure tensors.245 This is useful when you're only interested in the single trace of246 execution (e.g. for flop count), but not interested in247 generalizing the traced graph to new inputs.248 """249 super().__init__()250 if isinstance(model, (nn.parallel.distributed.DistributedDataParallel, nn.DataParallel)):251 model = model.module252 self.model = model253 if not isinstance(inputs, tuple):254 inputs = (inputs,)255 self.inputs = inputs256 self.allow_non_tensor = allow_non_tensor257 258 if inference_func is None:259 inference_func = lambda model, *inputs: model(*inputs) # noqa260 self.inference_func = inference_func261 262 self.flattened_inputs, self.inputs_schema = flatten_to_tuple(inputs)263 264 if all(isinstance(x, torch.Tensor) for x in self.flattened_inputs):265 return266 if self.allow_non_tensor:267 self.flattened_inputs = tuple(268 [x for x in self.flattened_inputs if isinstance(x, torch.Tensor)]269 )270 self.inputs_schema = None271 else:272 for input in self.flattened_inputs:273 if not isinstance(input, torch.Tensor):274 raise ValueError(275 "Inputs for tracing must only contain tensors. "276 f"Got a {type(input)} instead."277 )278 279 def forward(self, *args: torch.Tensor):280 with torch.no_grad(), patch_builtin_len():281 if self.inputs_schema is not None:282 inputs_orig_format = self.inputs_schema(args)283 else:284 if len(args) != len(self.flattened_inputs) or any(285 x is not y for x, y in zip(args, self.flattened_inputs)286 ):287 raise ValueError(288 "TracingAdapter does not contain valid inputs_schema."289 " So it cannot generalize to other inputs and must be"290 " traced with `.flattened_inputs`."291 )292 inputs_orig_format = self.inputs293 294 outputs = self.inference_func(self.model, *inputs_orig_format)295 flattened_outputs, schema = flatten_to_tuple(outputs)296 297 flattened_output_tensors = tuple(298 [x for x in flattened_outputs if isinstance(x, torch.Tensor)]299 )300 if len(flattened_output_tensors) < len(flattened_outputs):301 if self.allow_non_tensor:302 flattened_outputs = flattened_output_tensors303 self.outputs_schema = None304 else:305 raise ValueError(306 "Model cannot be traced because some model outputs "307 "cannot flatten to tensors."308 )309 else: # schema is valid310 if self.outputs_schema is None:311 self.outputs_schema = schema312 else:313 assert self.outputs_schema == schema, (314 "Model should always return outputs with the same "315 "structure so it can be traced!"316 )317 return flattened_outputs318 319 def _create_wrapper(self, traced_model):320 """321 Return a function that has an input/output interface the same as the322 original model, but it calls the given traced model under the hood.323 """324 325 def forward(*args):326 flattened_inputs, _ = flatten_to_tuple(args)327 flattened_outputs = traced_model(*flattened_inputs)328 return self.outputs_schema(flattened_outputs)329 330 return forward331 