Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
compat.py230 linesDownload Raw Back to config
1# Copyright (c) Facebook, Inc. and its affiliates.2"""3Backward compatibility of configs.4 5Instructions to bump version:6+ It's not needed to bump version if new keys are added.7  It's only needed when backward-incompatible changes happen8  (i.e., some existing keys disappear, or the meaning of a key changes)9+ To bump version, do the following:10    1. Increment _C.VERSION in defaults.py11    2. Add a converter in this file.12 13      Each ConverterVX has a function "upgrade" which in-place upgrades config from X-1 to X,14      and a function "downgrade" which in-place downgrades config from X to X-115 16      In each function, VERSION is left unchanged.17 18      Each converter assumes that its input has the relevant keys19      (i.e., the input is not a partial config).20    3. Run the tests (test_config.py) to make sure the upgrade & downgrade21       functions are consistent.22"""23 24import logging25from typing import List, Optional, Tuple26 27from .config import CfgNode as CN28from .defaults import _C29 30__all__ = ["upgrade_config", "downgrade_config"]31 32 33def upgrade_config(cfg: CN, to_version: Optional[int] = None) -> CN:34    """35    Upgrade a config from its current version to a newer version.36 37    Args:38        cfg (CfgNode):39        to_version (int): defaults to the latest version.40    """41    cfg = cfg.clone()42    if to_version is None:43        to_version = _C.VERSION44 45    assert cfg.VERSION <= to_version, "Cannot upgrade from v{} to v{}!".format(46        cfg.VERSION, to_version47    )48    for k in range(cfg.VERSION, to_version):49        converter = globals()["ConverterV" + str(k + 1)]50        converter.upgrade(cfg)51        cfg.VERSION = k + 152    return cfg53 54 55def downgrade_config(cfg: CN, to_version: int) -> CN:56    """57    Downgrade a config from its current version to an older version.58 59    Args:60        cfg (CfgNode):61        to_version (int):62 63    Note:64        A general downgrade of arbitrary configs is not always possible due to the65        different functionalities in different versions.66        The purpose of downgrade is only to recover the defaults in old versions,67        allowing it to load an old partial yaml config.68        Therefore, the implementation only needs to fill in the default values69        in the old version when a general downgrade is not possible.70    """71    cfg = cfg.clone()72    assert cfg.VERSION >= to_version, "Cannot downgrade from v{} to v{}!".format(73        cfg.VERSION, to_version74    )75    for k in range(cfg.VERSION, to_version, -1):76        converter = globals()["ConverterV" + str(k)]77        converter.downgrade(cfg)78        cfg.VERSION = k - 179    return cfg80 81 82def guess_version(cfg: CN, filename: str) -> int:83    """84    Guess the version of a partial config where the VERSION field is not specified.85    Returns the version, or the latest if cannot make a guess.86 87    This makes it easier for users to migrate.88    """89    logger = logging.getLogger(__name__)90 91    def _has(name: str) -> bool:92        cur = cfg93        for n in name.split("."):94            if n not in cur:95                return False96            cur = cur[n]97        return True98 99    # Most users' partial configs have "MODEL.WEIGHT", so guess on it100    ret = None101    if _has("MODEL.WEIGHT") or _has("TEST.AUG_ON"):102        ret = 1103 104    if ret is not None:105        logger.warning("Config '{}' has no VERSION. Assuming it to be v{}.".format(filename, ret))106    else:107        ret = _C.VERSION108        logger.warning(109            "Config '{}' has no VERSION. Assuming it to be compatible with latest v{}.".format(110                filename, ret111            )112        )113    return ret114 115 116def _rename(cfg: CN, old: str, new: str) -> None:117    old_keys = old.split(".")118    new_keys = new.split(".")119 120    def _set(key_seq: List[str], val: str) -> None:121        cur = cfg122        for k in key_seq[:-1]:123            if k not in cur:124                cur[k] = CN()125            cur = cur[k]126        cur[key_seq[-1]] = val127 128    def _get(key_seq: List[str]) -> CN:129        cur = cfg130        for k in key_seq:131            cur = cur[k]132        return cur133 134    def _del(key_seq: List[str]) -> None:135        cur = cfg136        for k in key_seq[:-1]:137            cur = cur[k]138        del cur[key_seq[-1]]139        if len(cur) == 0 and len(key_seq) > 1:140            _del(key_seq[:-1])141 142    _set(new_keys, _get(old_keys))143    _del(old_keys)144 145 146class _RenameConverter:147    """148    A converter that handles simple rename.149    """150 151    RENAME: List[Tuple[str, str]] = []  # list of tuples of (old name, new name)152 153    @classmethod154    def upgrade(cls, cfg: CN) -> None:155        for old, new in cls.RENAME:156            _rename(cfg, old, new)157 158    @classmethod159    def downgrade(cls, cfg: CN) -> None:160        for old, new in cls.RENAME[::-1]:161            _rename(cfg, new, old)162 163 164class ConverterV1(_RenameConverter):165    RENAME = [("MODEL.RPN_HEAD.NAME", "MODEL.RPN.HEAD_NAME")]166 167 168class ConverterV2(_RenameConverter):169    """170    A large bulk of rename, before public release.171    """172 173    RENAME = [174        ("MODEL.WEIGHT", "MODEL.WEIGHTS"),175        ("MODEL.PANOPTIC_FPN.SEMANTIC_LOSS_SCALE", "MODEL.SEM_SEG_HEAD.LOSS_WEIGHT"),176        ("MODEL.PANOPTIC_FPN.RPN_LOSS_SCALE", "MODEL.RPN.LOSS_WEIGHT"),177        ("MODEL.PANOPTIC_FPN.INSTANCE_LOSS_SCALE", "MODEL.PANOPTIC_FPN.INSTANCE_LOSS_WEIGHT"),178        ("MODEL.PANOPTIC_FPN.COMBINE_ON", "MODEL.PANOPTIC_FPN.COMBINE.ENABLED"),179        (180            "MODEL.PANOPTIC_FPN.COMBINE_OVERLAP_THRESHOLD",181            "MODEL.PANOPTIC_FPN.COMBINE.OVERLAP_THRESH",182        ),183        (184            "MODEL.PANOPTIC_FPN.COMBINE_STUFF_AREA_LIMIT",185            "MODEL.PANOPTIC_FPN.COMBINE.STUFF_AREA_LIMIT",186        ),187        (188            "MODEL.PANOPTIC_FPN.COMBINE_INSTANCES_CONFIDENCE_THRESHOLD",189            "MODEL.PANOPTIC_FPN.COMBINE.INSTANCES_CONFIDENCE_THRESH",190        ),191        ("MODEL.ROI_HEADS.SCORE_THRESH", "MODEL.ROI_HEADS.SCORE_THRESH_TEST"),192        ("MODEL.ROI_HEADS.NMS", "MODEL.ROI_HEADS.NMS_THRESH_TEST"),193        ("MODEL.RETINANET.INFERENCE_SCORE_THRESHOLD", "MODEL.RETINANET.SCORE_THRESH_TEST"),194        ("MODEL.RETINANET.INFERENCE_TOPK_CANDIDATES", "MODEL.RETINANET.TOPK_CANDIDATES_TEST"),195        ("MODEL.RETINANET.INFERENCE_NMS_THRESHOLD", "MODEL.RETINANET.NMS_THRESH_TEST"),196        ("TEST.DETECTIONS_PER_IMG", "TEST.DETECTIONS_PER_IMAGE"),197        ("TEST.AUG_ON", "TEST.AUG.ENABLED"),198        ("TEST.AUG_MIN_SIZES", "TEST.AUG.MIN_SIZES"),199        ("TEST.AUG_MAX_SIZE", "TEST.AUG.MAX_SIZE"),200        ("TEST.AUG_FLIP", "TEST.AUG.FLIP"),201    ]202 203    @classmethod204    def upgrade(cls, cfg: CN) -> None:205        super().upgrade(cfg)206 207        if cfg.MODEL.META_ARCHITECTURE == "RetinaNet":208            _rename(209                cfg, "MODEL.RETINANET.ANCHOR_ASPECT_RATIOS", "MODEL.ANCHOR_GENERATOR.ASPECT_RATIOS"210            )211            _rename(cfg, "MODEL.RETINANET.ANCHOR_SIZES", "MODEL.ANCHOR_GENERATOR.SIZES")212            del cfg["MODEL"]["RPN"]["ANCHOR_SIZES"]213            del cfg["MODEL"]["RPN"]["ANCHOR_ASPECT_RATIOS"]214        else:215            _rename(cfg, "MODEL.RPN.ANCHOR_ASPECT_RATIOS", "MODEL.ANCHOR_GENERATOR.ASPECT_RATIOS")216            _rename(cfg, "MODEL.RPN.ANCHOR_SIZES", "MODEL.ANCHOR_GENERATOR.SIZES")217            del cfg["MODEL"]["RETINANET"]["ANCHOR_SIZES"]218            del cfg["MODEL"]["RETINANET"]["ANCHOR_ASPECT_RATIOS"]219        del cfg["MODEL"]["RETINANET"]["ANCHOR_STRIDES"]220 221    @classmethod222    def downgrade(cls, cfg: CN) -> None:223        super().downgrade(cfg)224 225        _rename(cfg, "MODEL.ANCHOR_GENERATOR.ASPECT_RATIOS", "MODEL.RPN.ANCHOR_ASPECT_RATIOS")226        _rename(cfg, "MODEL.ANCHOR_GENERATOR.SIZES", "MODEL.RPN.ANCHOR_SIZES")227        cfg.MODEL.RETINANET.ANCHOR_ASPECT_RATIOS = cfg.MODEL.RPN.ANCHOR_ASPECT_RATIOS228        cfg.MODEL.RETINANET.ANCHOR_SIZES = cfg.MODEL.RPN.ANCHOR_SIZES229        cfg.MODEL.RETINANET.ANCHOR_STRIDES = []  # this is not used anywhere in any version230