Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2import logging3import os4import pickle5from urllib.parse import parse_qs, urlparse6import torch7from fvcore.common.checkpoint import Checkpointer8from torch.nn.parallel import DistributedDataParallel9 10import detectron2.utils.comm as comm11from detectron2.utils.file_io import PathManager12 13from .c2_model_loading import align_and_update_state_dicts14 15 16class DetectionCheckpointer(Checkpointer):17 """18 Same as :class:`Checkpointer`, but is able to:19 1. handle models in detectron & detectron2 model zoo, and apply conversions for legacy models.20 2. correctly load checkpoints that are only available on the master worker21 """22 23 def __init__(self, model, save_dir="", *, save_to_disk=None, **checkpointables):24 is_main_process = comm.is_main_process()25 super().__init__(26 model,27 save_dir,28 save_to_disk=is_main_process if save_to_disk is None else save_to_disk,29 **checkpointables,30 )31 self.path_manager = PathManager32 self._parsed_url_during_load = None33 34 def load(self, path, *args, **kwargs):35 assert self._parsed_url_during_load is None36 need_sync = False37 logger = logging.getLogger(__name__)38 logger.info("[DetectionCheckpointer] Loading from {} ...".format(path))39 40 if path and isinstance(self.model, DistributedDataParallel):41 path = self.path_manager.get_local_path(path)42 has_file = os.path.isfile(path)43 all_has_file = comm.all_gather(has_file)44 if not all_has_file[0]:45 raise OSError(f"File {path} not found on main worker.")46 if not all(all_has_file):47 logger.warning(48 f"Not all workers can read checkpoint {path}. "49 "Training may fail to fully resume."50 )51 # TODO: broadcast the checkpoint file contents from main52 # worker, and load from it instead.53 need_sync = True54 if not has_file:55 path = None # don't load if not readable56 57 if path:58 parsed_url = urlparse(path)59 self._parsed_url_during_load = parsed_url60 path = parsed_url._replace(query="").geturl() # remove query from filename61 path = self.path_manager.get_local_path(path)62 ret = super().load(path, *args, **kwargs)63 64 if need_sync:65 logger.info("Broadcasting model states from main worker ...")66 self.model._sync_params_and_buffers()67 self._parsed_url_during_load = None # reset to None68 return ret69 70 def _load_file(self, filename):71 if filename.endswith(".pkl"):72 with PathManager.open(filename, "rb") as f:73 data = pickle.load(f, encoding="latin1")74 if "model" in data and "__author__" in data:75 # file is in Detectron2 model zoo format76 self.logger.info("Reading a file from '{}'".format(data["__author__"]))77 return data78 else:79 # assume file is from Caffe2 / Detectron1 model zoo80 if "blobs" in data:81 # Detection models have "blobs", but ImageNet models don't82 data = data["blobs"]83 data = {k: v for k, v in data.items() if not k.endswith("_momentum")}84 return {"model": data, "__author__": "Caffe2", "matching_heuristics": True}85 elif filename.endswith(".pyth"):86 # assume file is from pycls; no one else seems to use the ".pyth" extension87 with PathManager.open(filename, "rb") as f:88 data = torch.load(f)89 assert (90 "model_state" in data91 ), f"Cannot load .pyth file {filename}; pycls checkpoints must contain 'model_state'."92 model_state = {93 k: v94 for k, v in data["model_state"].items()95 if not k.endswith("num_batches_tracked")96 }97 return {"model": model_state, "__author__": "pycls", "matching_heuristics": True}98 99 loaded = self._torch_load(filename)100 if "model" not in loaded:101 loaded = {"model": loaded}102 assert self._parsed_url_during_load is not None, "`_load_file` must be called inside `load`"103 parsed_url = self._parsed_url_during_load104 queries = parse_qs(parsed_url.query)105 if queries.pop("matching_heuristics", "False") == ["True"]:106 loaded["matching_heuristics"] = True107 if len(queries) > 0:108 raise ValueError(109 f"Unsupported query remaining: f{queries}, orginal filename: {parsed_url.geturl()}"110 )111 return loaded112 113 def _torch_load(self, f):114 return super()._load_file(f)115 116 def _load_model(self, checkpoint):117 if checkpoint.get("matching_heuristics", False):118 self._convert_ndarray_to_tensor(checkpoint["model"])119 # convert weights by name-matching heuristics120 checkpoint["model"] = align_and_update_state_dicts(121 self.model.state_dict(),122 checkpoint["model"],123 c2_conversion=checkpoint.get("__author__", None) == "Caffe2",124 )125 # for non-caffe2 models, use standard ways to load it126 incompatible = super()._load_model(checkpoint)127 128 model_buffers = dict(self.model.named_buffers(recurse=False))129 for k in ["pixel_mean", "pixel_std"]:130 # Ignore missing key message about pixel_mean/std.131 # Though they may be missing in old checkpoints, they will be correctly132 # initialized from config anyway.133 if k in model_buffers:134 try:135 incompatible.missing_keys.remove(k)136 except ValueError:137 pass138 for k in incompatible.unexpected_keys[:]:139 # Ignore unexpected keys about cell anchors. They exist in old checkpoints140 # but now they are non-persistent buffers and will not be in new checkpoints.141 if "anchor_generator.cell_anchors" in k:142 incompatible.unexpected_keys.remove(k)143 return incompatible144 