Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2 3import os4import sys5import tempfile6from contextlib import ExitStack, contextmanager7from copy import deepcopy8from unittest import mock9import torch10from torch import nn11 12# need some explicit imports due to https://github.com/pytorch/pytorch/issues/3896413import detectron2 # noqa F40114from detectron2.structures import Boxes, Instances15from detectron2.utils.env import _import_file16 17_counter = 018 19 20def _clear_jit_cache():21 from torch.jit._recursive import concrete_type_store22 from torch.jit._state import _jit_caching_layer23 24 concrete_type_store.type_store.clear() # for modules25 _jit_caching_layer.clear() # for free functions26 27 28def _add_instances_conversion_methods(newInstances):29 """30 Add from_instances methods to the scripted Instances class.31 """32 cls_name = newInstances.__name__33 34 @torch.jit.unused35 def from_instances(instances: Instances):36 """37 Create scripted Instances from original Instances38 """39 fields = instances.get_fields()40 image_size = instances.image_size41 ret = newInstances(image_size)42 for name, val in fields.items():43 assert hasattr(ret, f"_{name}"), f"No attribute named {name} in {cls_name}"44 setattr(ret, name, deepcopy(val))45 return ret46 47 newInstances.from_instances = from_instances48 49 50@contextmanager51def patch_instances(fields):52 """53 A contextmanager, under which the Instances class in detectron2 is replaced54 by a statically-typed scriptable class, defined by `fields`.55 See more in `scripting_with_instances`.56 """57 58 with tempfile.TemporaryDirectory(prefix="detectron2") as dir, tempfile.NamedTemporaryFile(59 mode="w", encoding="utf-8", suffix=".py", dir=dir, delete=False60 ) as f:61 try:62 # Objects that use Instances should not reuse previously-compiled63 # results in cache, because `Instances` could be a new class each time.64 _clear_jit_cache()65 66 cls_name, s = _gen_instance_module(fields)67 f.write(s)68 f.flush()69 f.close()70 71 module = _import(f.name)72 new_instances = getattr(module, cls_name)73 _ = torch.jit.script(new_instances)74 # let torchscript think Instances was scripted already75 Instances.__torch_script_class__ = True76 # let torchscript find new_instances when looking for the jit type of Instances77 Instances._jit_override_qualname = torch._jit_internal._qualified_name(new_instances)78 79 _add_instances_conversion_methods(new_instances)80 yield new_instances81 finally:82 try:83 del Instances.__torch_script_class__84 del Instances._jit_override_qualname85 except AttributeError:86 pass87 sys.modules.pop(module.__name__)88 89 90def _gen_instance_class(fields):91 """92 Args:93 fields (dict[name: type])94 """95 96 class _FieldType:97 def __init__(self, name, type_):98 assert isinstance(name, str), f"Field name must be str, got {name}"99 self.name = name100 self.type_ = type_101 self.annotation = f"{type_.__module__}.{type_.__name__}"102 103 fields = [_FieldType(k, v) for k, v in fields.items()]104 105 def indent(level, s):106 return " " * 4 * level + s107 108 lines = []109 110 global _counter111 _counter += 1112 113 cls_name = "ScriptedInstances{}".format(_counter)114 115 field_names = tuple(x.name for x in fields)116 extra_args = ", ".join([f"{f.name}: Optional[{f.annotation}] = None" for f in fields])117 lines.append(118 f"""119class {cls_name}:120 def __init__(self, image_size: Tuple[int, int], {extra_args}):121 self.image_size = image_size122 self._field_names = {field_names}123"""124 )125 126 for f in fields:127 lines.append(128 indent(2, f"self._{f.name} = torch.jit.annotate(Optional[{f.annotation}], {f.name})")129 )130 131 for f in fields:132 lines.append(133 f"""134 @property135 def {f.name}(self) -> {f.annotation}:136 # has to use a local for type refinement137 # https://pytorch.org/docs/stable/jit_language_reference.html#optional-type-refinement138 t = self._{f.name}139 assert t is not None, "{f.name} is None and cannot be accessed!"140 return t141 142 @{f.name}.setter143 def {f.name}(self, value: {f.annotation}) -> None:144 self._{f.name} = value145"""146 )147 148 # support method `__len__`149 lines.append(150 """151 def __len__(self) -> int:152"""153 )154 for f in fields:155 lines.append(156 f"""157 t = self._{f.name}158 if t is not None:159 return len(t)160"""161 )162 lines.append(163 """164 raise NotImplementedError("Empty Instances does not support __len__!")165"""166 )167 168 # support method `has`169 lines.append(170 """171 def has(self, name: str) -> bool:172"""173 )174 for f in fields:175 lines.append(176 f"""177 if name == "{f.name}":178 return self._{f.name} is not None179"""180 )181 lines.append(182 """183 return False184"""185 )186 187 # support method `to`188 none_args = ", None" * len(fields)189 lines.append(190 f"""191 def to(self, device: torch.device) -> "{cls_name}":192 ret = {cls_name}(self.image_size{none_args})193"""194 )195 for f in fields:196 if hasattr(f.type_, "to"):197 lines.append(198 f"""199 t = self._{f.name}200 if t is not None:201 ret._{f.name} = t.to(device)202"""203 )204 else:205 # For now, ignore fields that cannot be moved to devices.206 # Maybe can support other tensor-like classes (e.g. __torch_function__)207 pass208 lines.append(209 """210 return ret211"""212 )213 214 # support method `getitem`215 none_args = ", None" * len(fields)216 lines.append(217 f"""218 def __getitem__(self, item) -> "{cls_name}":219 ret = {cls_name}(self.image_size{none_args})220"""221 )222 for f in fields:223 lines.append(224 f"""225 t = self._{f.name}226 if t is not None:227 ret._{f.name} = t[item]228"""229 )230 lines.append(231 """232 return ret233"""234 )235 236 # support method `cat`237 # this version does not contain checks that all instances have same size and fields238 none_args = ", None" * len(fields)239 lines.append(240 f"""241 def cat(self, instances: List["{cls_name}"]) -> "{cls_name}":242 ret = {cls_name}(self.image_size{none_args})243"""244 )245 for f in fields:246 lines.append(247 f"""248 t = self._{f.name}249 if t is not None:250 values: List[{f.annotation}] = [x.{f.name} for x in instances]251 if torch.jit.isinstance(t, torch.Tensor):252 ret._{f.name} = torch.cat(values, dim=0)253 else:254 ret._{f.name} = t.cat(values)255"""256 )257 lines.append(258 """259 return ret"""260 )261 262 # support method `get_fields()`263 lines.append(264 """265 def get_fields(self) -> Dict[str, Tensor]:266 ret = {}267 """268 )269 for f in fields:270 if f.type_ == Boxes:271 stmt = "t.tensor"272 elif f.type_ == torch.Tensor:273 stmt = "t"274 else:275 stmt = f'assert False, "unsupported type {str(f.type_)}"'276 lines.append(277 f"""278 t = self._{f.name}279 if t is not None:280 ret["{f.name}"] = {stmt}281 """282 )283 lines.append(284 """285 return ret"""286 )287 return cls_name, os.linesep.join(lines)288 289 290def _gen_instance_module(fields):291 # TODO: find a more automatic way to enable import of other classes292 s = """293from copy import deepcopy294import torch295from torch import Tensor296import typing297from typing import *298 299import detectron2300from detectron2.structures import Boxes, Instances301 302"""303 304 cls_name, cls_def = _gen_instance_class(fields)305 s += cls_def306 return cls_name, s307 308 309def _import(path):310 return _import_file(311 "{}{}".format(sys.modules[__name__].__name__, _counter), path, make_importable=True312 )313 314 315@contextmanager316def patch_builtin_len(modules=()):317 """318 Patch the builtin len() function of a few detectron2 modules319 to use __len__ instead, because __len__ does not convert values to320 integers and therefore is friendly to tracing.321 322 Args:323 modules (list[stsr]): names of extra modules to patch len(), in324 addition to those in detectron2.325 """326 327 def _new_len(obj):328 return obj.__len__()329 330 with ExitStack() as stack:331 MODULES = [332 "detectron2.modeling.roi_heads.fast_rcnn",333 "detectron2.modeling.roi_heads.mask_head",334 "detectron2.modeling.roi_heads.keypoint_head",335 ] + list(modules)336 ctxs = [stack.enter_context(mock.patch(mod + ".len")) for mod in MODULES]337 for m in ctxs:338 m.side_effect = _new_len339 yield340 341 342def patch_nonscriptable_classes():343 """344 Apply patches on a few nonscriptable detectron2 classes.345 Should not have side-effects on eager usage.346 """347 # __prepare_scriptable__ can also be added to models for easier maintenance.348 # But it complicates the clean model code.349 350 from detectron2.modeling.backbone import ResNet, FPN351 352 # Due to https://github.com/pytorch/pytorch/issues/36061,353 # we change backbone to use ModuleList for scripting.354 # (note: this changes param names in state_dict)355 356 def prepare_resnet(self):357 ret = deepcopy(self)358 ret.stages = nn.ModuleList(ret.stages)359 for k in self.stage_names:360 delattr(ret, k)361 return ret362 363 ResNet.__prepare_scriptable__ = prepare_resnet364 365 def prepare_fpn(self):366 ret = deepcopy(self)367 ret.lateral_convs = nn.ModuleList(ret.lateral_convs)368 ret.output_convs = nn.ModuleList(ret.output_convs)369 for name, _ in self.named_children():370 if name.startswith("fpn_"):371 delattr(ret, name)372 return ret373 374 FPN.__prepare_scriptable__ = prepare_fpn375 376 # Annotate some attributes to be constants for the purpose of scripting,377 # even though they are not constants in eager mode.378 from detectron2.modeling.roi_heads import StandardROIHeads379 380 if hasattr(StandardROIHeads, "__annotations__"):381 # copy first to avoid editing annotations of base class382 StandardROIHeads.__annotations__ = deepcopy(StandardROIHeads.__annotations__)383 StandardROIHeads.__annotations__["mask_on"] = torch.jit.Final[bool]384 StandardROIHeads.__annotations__["keypoint_on"] = torch.jit.Final[bool]385 386 387# These patches are not supposed to have side-effects.388patch_nonscriptable_classes()389 390 391@contextmanager392def freeze_training_mode(model):393 """394 A context manager that annotates the "training" attribute of every submodule395 to constant, so that the training codepath in these modules can be396 meta-compiled away. Upon exiting, the annotations are reverted.397 """398 classes = {type(x) for x in model.modules()}399 # __constants__ is the old way to annotate constants and not compatible400 # with __annotations__ .401 classes = {x for x in classes if not hasattr(x, "__constants__")}402 for cls in classes:403 cls.__annotations__["training"] = torch.jit.Final[bool]404 yield405 for cls in classes:406 cls.__annotations__["training"] = bool407 