Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
detection_checkpoint.py144 linesDownload Raw Back to checkpoint
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