Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
torchscript.py133 linesDownload Raw Back to export
1# Copyright (c) Facebook, Inc. and its affiliates.2 3import os4import torch5 6from detectron2.utils.file_io import PathManager7 8from .torchscript_patch import freeze_training_mode, patch_instances9 10__all__ = ["scripting_with_instances", "dump_torchscript_IR"]11 12 13def scripting_with_instances(model, fields):14    """15    Run :func:`torch.jit.script` on a model that uses the :class:`Instances` class. Since16    attributes of :class:`Instances` are "dynamically" added in eager mode,it is difficult17    for scripting to support it out of the box. This function is made to support scripting18    a model that uses :class:`Instances`. It does the following:19 20    1. Create a scriptable ``new_Instances`` class which behaves similarly to ``Instances``,21       but with all attributes been "static".22       The attributes need to be statically declared in the ``fields`` argument.23    2. Register ``new_Instances``, and force scripting compiler to24       use it when trying to compile ``Instances``.25 26    After this function, the process will be reverted. User should be able to script another model27    using different fields.28 29    Example:30        Assume that ``Instances`` in the model consist of two attributes named31        ``proposal_boxes`` and ``objectness_logits`` with type :class:`Boxes` and32        :class:`Tensor` respectively during inference. You can call this function like:33        ::34            fields = {"proposal_boxes": Boxes, "objectness_logits": torch.Tensor}35            torchscipt_model =  scripting_with_instances(model, fields)36 37    Note:38        It only support models in evaluation mode.39 40    Args:41        model (nn.Module): The input model to be exported by scripting.42        fields (Dict[str, type]): Attribute names and corresponding type that43            ``Instances`` will use in the model. Note that all attributes used in ``Instances``44            need to be added, regardless of whether they are inputs/outputs of the model.45            Data type not defined in detectron2 is not supported for now.46 47    Returns:48        torch.jit.ScriptModule: the model in torchscript format49    """50    assert (51        not model.training52    ), "Currently we only support exporting models in evaluation mode to torchscript"53 54    with freeze_training_mode(model), patch_instances(fields):55        scripted_model = torch.jit.script(model)56        return scripted_model57 58 59# alias for old name60export_torchscript_with_instances = scripting_with_instances61 62 63def dump_torchscript_IR(model, dir):64    """65    Dump IR of a TracedModule/ScriptModule/Function in various format (code, graph,66    inlined graph). Useful for debugging.67 68    Args:69        model (TracedModule/ScriptModule/ScriptFUnction): traced or scripted module70        dir (str): output directory to dump files.71    """72    dir = os.path.expanduser(dir)73    PathManager.mkdirs(dir)74 75    def _get_script_mod(mod):76        if isinstance(mod, torch.jit.TracedModule):77            return mod._actual_script_module78        return mod79 80    # Dump pretty-printed code: https://pytorch.org/docs/stable/jit.html#inspecting-code81    with PathManager.open(os.path.join(dir, "model_ts_code.txt"), "w") as f:82 83        def get_code(mod):84            # Try a few ways to get code using private attributes.85            try:86                # This contains more information than just `mod.code`87                return _get_script_mod(mod)._c.code88            except AttributeError:89                pass90            try:91                return mod.code92            except AttributeError:93                return None94 95        def dump_code(prefix, mod):96            code = get_code(mod)97            name = prefix or "root model"98            if code is None:99                f.write(f"Could not found code for {name} (type={mod.original_name})\n")100                f.write("\n")101            else:102                f.write(f"\nCode for {name}, type={mod.original_name}:\n")103                f.write(code)104                f.write("\n")105                f.write("-" * 80)106 107            for name, m in mod.named_children():108                dump_code(prefix + "." + name, m)109 110        if isinstance(model, torch.jit.ScriptFunction):111            f.write(get_code(model))112        else:113            dump_code("", model)114 115    def _get_graph(model):116        try:117            # Recursively dump IR of all modules118            return _get_script_mod(model)._c.dump_to_str(True, False, False)119        except AttributeError:120            return model.graph.str()121 122    with PathManager.open(os.path.join(dir, "model_ts_IR.txt"), "w") as f:123        f.write(_get_graph(model))124 125    # Dump IR of the entire graph (all submodules inlined)126    with PathManager.open(os.path.join(dir, "model_ts_IR_inlined.txt"), "w") as f:127        f.write(str(model.inlined_graph))128 129    if not isinstance(model, torch.jit.ScriptFunction):130        # Dump the model structure in pytorch style131        with PathManager.open(os.path.join(dir, "model.txt"), "w") as f:132            f.write(str(model))133