Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
torchscript_patch.py407 linesDownload Raw Back to export
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