Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
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 