Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

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