Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
c2_model_loading.py413 linesDownload Raw Back to checkpoint
1# Copyright (c) Facebook, Inc. and its affiliates.2import copy3import logging4import re5from typing import Dict, List6import torch7from tabulate import tabulate8 9 10def convert_basic_c2_names(original_keys):11    """12    Apply some basic name conversion to names in C2 weights.13    It only deals with typical backbone models.14 15    Args:16        original_keys (list[str]):17    Returns:18        list[str]: The same number of strings matching those in original_keys.19    """20    layer_keys = copy.deepcopy(original_keys)21    layer_keys = [22        {"pred_b": "linear_b", "pred_w": "linear_w"}.get(k, k) for k in layer_keys23    ]  # some hard-coded mappings24 25    layer_keys = [k.replace("_", ".") for k in layer_keys]26    layer_keys = [re.sub("\\.b$", ".bias", k) for k in layer_keys]27    layer_keys = [re.sub("\\.w$", ".weight", k) for k in layer_keys]28    # Uniform both bn and gn names to "norm"29    layer_keys = [re.sub("bn\\.s$", "norm.weight", k) for k in layer_keys]30    layer_keys = [re.sub("bn\\.bias$", "norm.bias", k) for k in layer_keys]31    layer_keys = [re.sub("bn\\.rm", "norm.running_mean", k) for k in layer_keys]32    layer_keys = [re.sub("bn\\.running.mean$", "norm.running_mean", k) for k in layer_keys]33    layer_keys = [re.sub("bn\\.riv$", "norm.running_var", k) for k in layer_keys]34    layer_keys = [re.sub("bn\\.running.var$", "norm.running_var", k) for k in layer_keys]35    layer_keys = [re.sub("bn\\.gamma$", "norm.weight", k) for k in layer_keys]36    layer_keys = [re.sub("bn\\.beta$", "norm.bias", k) for k in layer_keys]37    layer_keys = [re.sub("gn\\.s$", "norm.weight", k) for k in layer_keys]38    layer_keys = [re.sub("gn\\.bias$", "norm.bias", k) for k in layer_keys]39 40    # stem41    layer_keys = [re.sub("^res\\.conv1\\.norm\\.", "conv1.norm.", k) for k in layer_keys]42    # to avoid mis-matching with "conv1" in other components (e.g. detection head)43    layer_keys = [re.sub("^conv1\\.", "stem.conv1.", k) for k in layer_keys]44 45    # layer1-4 is used by torchvision, however we follow the C2 naming strategy (res2-5)46    # layer_keys = [re.sub("^res2.", "layer1.", k) for k in layer_keys]47    # layer_keys = [re.sub("^res3.", "layer2.", k) for k in layer_keys]48    # layer_keys = [re.sub("^res4.", "layer3.", k) for k in layer_keys]49    # layer_keys = [re.sub("^res5.", "layer4.", k) for k in layer_keys]50 51    # blocks52    layer_keys = [k.replace(".branch1.", ".shortcut.") for k in layer_keys]53    layer_keys = [k.replace(".branch2a.", ".conv1.") for k in layer_keys]54    layer_keys = [k.replace(".branch2b.", ".conv2.") for k in layer_keys]55    layer_keys = [k.replace(".branch2c.", ".conv3.") for k in layer_keys]56 57    # DensePose substitutions58    layer_keys = [re.sub("^body.conv.fcn", "body_conv_fcn", k) for k in layer_keys]59    layer_keys = [k.replace("AnnIndex.lowres", "ann_index_lowres") for k in layer_keys]60    layer_keys = [k.replace("Index.UV.lowres", "index_uv_lowres") for k in layer_keys]61    layer_keys = [k.replace("U.lowres", "u_lowres") for k in layer_keys]62    layer_keys = [k.replace("V.lowres", "v_lowres") for k in layer_keys]63    return layer_keys64 65 66def convert_c2_detectron_names(weights):67    """68    Map Caffe2 Detectron weight names to Detectron2 names.69 70    Args:71        weights (dict): name -> tensor72 73    Returns:74        dict: detectron2 names -> tensor75        dict: detectron2 names -> C2 names76    """77    logger = logging.getLogger(__name__)78    logger.info("Renaming Caffe2 weights ......")79    original_keys = sorted(weights.keys())80    layer_keys = copy.deepcopy(original_keys)81 82    layer_keys = convert_basic_c2_names(layer_keys)83 84    # --------------------------------------------------------------------------85    # RPN hidden representation conv86    # --------------------------------------------------------------------------87    # FPN case88    # In the C2 model, the RPN hidden layer conv is defined for FPN level 2 and then89    # shared for all other levels, hence the appearance of "fpn2"90    layer_keys = [91        k.replace("conv.rpn.fpn2", "proposal_generator.rpn_head.conv") for k in layer_keys92    ]93    # Non-FPN case94    layer_keys = [k.replace("conv.rpn", "proposal_generator.rpn_head.conv") for k in layer_keys]95 96    # --------------------------------------------------------------------------97    # RPN box transformation conv98    # --------------------------------------------------------------------------99    # FPN case (see note above about "fpn2")100    layer_keys = [101        k.replace("rpn.bbox.pred.fpn2", "proposal_generator.rpn_head.anchor_deltas")102        for k in layer_keys103    ]104    layer_keys = [105        k.replace("rpn.cls.logits.fpn2", "proposal_generator.rpn_head.objectness_logits")106        for k in layer_keys107    ]108    # Non-FPN case109    layer_keys = [110        k.replace("rpn.bbox.pred", "proposal_generator.rpn_head.anchor_deltas") for k in layer_keys111    ]112    layer_keys = [113        k.replace("rpn.cls.logits", "proposal_generator.rpn_head.objectness_logits")114        for k in layer_keys115    ]116 117    # --------------------------------------------------------------------------118    # Fast R-CNN box head119    # --------------------------------------------------------------------------120    layer_keys = [re.sub("^bbox\\.pred", "bbox_pred", k) for k in layer_keys]121    layer_keys = [re.sub("^cls\\.score", "cls_score", k) for k in layer_keys]122    layer_keys = [re.sub("^fc6\\.", "box_head.fc1.", k) for k in layer_keys]123    layer_keys = [re.sub("^fc7\\.", "box_head.fc2.", k) for k in layer_keys]124    # 4conv1fc head tensor names: head_conv1_w, head_conv1_gn_s125    layer_keys = [re.sub("^head\\.conv", "box_head.conv", k) for k in layer_keys]126 127    # --------------------------------------------------------------------------128    # FPN lateral and output convolutions129    # --------------------------------------------------------------------------130    def fpn_map(name):131        """132        Look for keys with the following patterns:133        1) Starts with "fpn.inner."134           Example: "fpn.inner.res2.2.sum.lateral.weight"135           Meaning: These are lateral pathway convolutions136        2) Starts with "fpn.res"137           Example: "fpn.res2.2.sum.weight"138           Meaning: These are FPN output convolutions139        """140        splits = name.split(".")141        norm = ".norm" if "norm" in splits else ""142        if name.startswith("fpn.inner."):143            # splits example: ['fpn', 'inner', 'res2', '2', 'sum', 'lateral', 'weight']144            stage = int(splits[2][len("res") :])145            return "fpn_lateral{}{}.{}".format(stage, norm, splits[-1])146        elif name.startswith("fpn.res"):147            # splits example: ['fpn', 'res2', '2', 'sum', 'weight']148            stage = int(splits[1][len("res") :])149            return "fpn_output{}{}.{}".format(stage, norm, splits[-1])150        return name151 152    layer_keys = [fpn_map(k) for k in layer_keys]153 154    # --------------------------------------------------------------------------155    # Mask R-CNN mask head156    # --------------------------------------------------------------------------157    # roi_heads.StandardROIHeads case158    layer_keys = [k.replace(".[mask].fcn", "mask_head.mask_fcn") for k in layer_keys]159    layer_keys = [re.sub("^\\.mask\\.fcn", "mask_head.mask_fcn", k) for k in layer_keys]160    layer_keys = [k.replace("mask.fcn.logits", "mask_head.predictor") for k in layer_keys]161    # roi_heads.Res5ROIHeads case162    layer_keys = [k.replace("conv5.mask", "mask_head.deconv") for k in layer_keys]163 164    # --------------------------------------------------------------------------165    # Keypoint R-CNN head166    # --------------------------------------------------------------------------167    # interestingly, the keypoint head convs have blob names that are simply "conv_fcnX"168    layer_keys = [k.replace("conv.fcn", "roi_heads.keypoint_head.conv_fcn") for k in layer_keys]169    layer_keys = [170        k.replace("kps.score.lowres", "roi_heads.keypoint_head.score_lowres") for k in layer_keys171    ]172    layer_keys = [k.replace("kps.score.", "roi_heads.keypoint_head.score.") for k in layer_keys]173 174    # --------------------------------------------------------------------------175    # Done with replacements176    # --------------------------------------------------------------------------177    assert len(set(layer_keys)) == len(layer_keys)178    assert len(original_keys) == len(layer_keys)179 180    new_weights = {}181    new_keys_to_original_keys = {}182    for orig, renamed in zip(original_keys, layer_keys):183        new_keys_to_original_keys[renamed] = orig184        if renamed.startswith("bbox_pred.") or renamed.startswith("mask_head.predictor."):185            # remove the meaningless prediction weight for background class186            new_start_idx = 4 if renamed.startswith("bbox_pred.") else 1187            new_weights[renamed] = weights[orig][new_start_idx:]188            logger.info(189                "Remove prediction weight for background class in {}. The shape changes from "190                "{} to {}.".format(191                    renamed, tuple(weights[orig].shape), tuple(new_weights[renamed].shape)192                )193            )194        elif renamed.startswith("cls_score."):195            # move weights of bg class from original index 0 to last index196            logger.info(197                "Move classification weights for background class in {} from index 0 to "198                "index {}.".format(renamed, weights[orig].shape[0] - 1)199            )200            new_weights[renamed] = torch.cat([weights[orig][1:], weights[orig][:1]])201        else:202            new_weights[renamed] = weights[orig]203 204    return new_weights, new_keys_to_original_keys205 206 207# Note the current matching is not symmetric.208# it assumes model_state_dict will have longer names.209def align_and_update_state_dicts(model_state_dict, ckpt_state_dict, c2_conversion=True):210    """211    Match names between the two state-dict, and returns a new chkpt_state_dict with names212    converted to match model_state_dict with heuristics. The returned dict can be later213    loaded with fvcore checkpointer.214    If `c2_conversion==True`, `ckpt_state_dict` is assumed to be a Caffe2215    model and will be renamed at first.216 217    Strategy: suppose that the models that we will create will have prefixes appended218    to each of its keys, for example due to an extra level of nesting that the original219    pre-trained weights from ImageNet won't contain. For example, model.state_dict()220    might return backbone[0].body.res2.conv1.weight, while the pre-trained model contains221    res2.conv1.weight. We thus want to match both parameters together.222    For that, we look for each model weight, look among all loaded keys if there is one223    that is a suffix of the current weight name, and use it if that's the case.224    If multiple matches exist, take the one with longest size225    of the corresponding name. For example, for the same model as before, the pretrained226    weight file can contain both res2.conv1.weight, as well as conv1.weight. In this case,227    we want to match backbone[0].body.conv1.weight to conv1.weight, and228    backbone[0].body.res2.conv1.weight to res2.conv1.weight.229    """230    model_keys = sorted(model_state_dict.keys())231    if c2_conversion:232        ckpt_state_dict, original_keys = convert_c2_detectron_names(ckpt_state_dict)233        # original_keys: the name in the original dict (before renaming)234    else:235        original_keys = {x: x for x in ckpt_state_dict.keys()}236    ckpt_keys = sorted(ckpt_state_dict.keys())237 238    def match(a, b):239        # Matched ckpt_key should be a complete (starts with '.') suffix.240        # For example, roi_heads.mesh_head.whatever_conv1 does not match conv1,241        # but matches whatever_conv1 or mesh_head.whatever_conv1.242        return a == b or a.endswith("." + b)243 244    # get a matrix of string matches, where each (i, j) entry correspond to the size of the245    # ckpt_key string, if it matches246    match_matrix = [len(j) if match(i, j) else 0 for i in model_keys for j in ckpt_keys]247    match_matrix = torch.as_tensor(match_matrix).view(len(model_keys), len(ckpt_keys))248    # use the matched one with longest size in case of multiple matches249    max_match_size, idxs = match_matrix.max(1)250    # remove indices that correspond to no-match251    idxs[max_match_size == 0] = -1252 253    logger = logging.getLogger(__name__)254    # matched_pairs (matched checkpoint key --> matched model key)255    matched_keys = {}256    result_state_dict = {}257    for idx_model, idx_ckpt in enumerate(idxs.tolist()):258        if idx_ckpt == -1:259            continue260        key_model = model_keys[idx_model]261        key_ckpt = ckpt_keys[idx_ckpt]262        value_ckpt = ckpt_state_dict[key_ckpt]263        shape_in_model = model_state_dict[key_model].shape264 265        if shape_in_model != value_ckpt.shape:266            logger.warning(267                "Shape of {} in checkpoint is {}, while shape of {} in model is {}.".format(268                    key_ckpt, value_ckpt.shape, key_model, shape_in_model269                )270            )271            logger.warning(272                "{} will not be loaded. Please double check and see if this is desired.".format(273                    key_ckpt274                )275            )276            continue277 278        assert key_model not in result_state_dict279        result_state_dict[key_model] = value_ckpt280        if key_ckpt in matched_keys:  # already added to matched_keys281            logger.error(282                "Ambiguity found for {} in checkpoint!"283                "It matches at least two keys in the model ({} and {}).".format(284                    key_ckpt, key_model, matched_keys[key_ckpt]285                )286            )287            raise ValueError("Cannot match one checkpoint key to multiple keys in the model.")288 289        matched_keys[key_ckpt] = key_model290 291    # logging:292    matched_model_keys = sorted(matched_keys.values())293    if len(matched_model_keys) == 0:294        logger.warning("No weights in checkpoint matched with model.")295        return ckpt_state_dict296    common_prefix = _longest_common_prefix(matched_model_keys)297    rev_matched_keys = {v: k for k, v in matched_keys.items()}298    original_keys = {k: original_keys[rev_matched_keys[k]] for k in matched_model_keys}299 300    model_key_groups = _group_keys_by_module(matched_model_keys, original_keys)301    table = []302    memo = set()303    for key_model in matched_model_keys:304        if key_model in memo:305            continue306        if key_model in model_key_groups:307            group = model_key_groups[key_model]308            memo |= set(group)309            shapes = [tuple(model_state_dict[k].shape) for k in group]310            table.append(311                (312                    _longest_common_prefix([k[len(common_prefix) :] for k in group]) + "*",313                    _group_str([original_keys[k] for k in group]),314                    " ".join([str(x).replace(" ", "") for x in shapes]),315                )316            )317        else:318            key_checkpoint = original_keys[key_model]319            shape = str(tuple(model_state_dict[key_model].shape))320            table.append((key_model[len(common_prefix) :], key_checkpoint, shape))321    table_str = tabulate(322        table, tablefmt="pipe", headers=["Names in Model", "Names in Checkpoint", "Shapes"]323    )324    logger.info(325        "Following weights matched with "326        + (f"submodule {common_prefix[:-1]}" if common_prefix else "model")327        + ":\n"328        + table_str329    )330 331    unmatched_ckpt_keys = [k for k in ckpt_keys if k not in set(matched_keys.keys())]332    for k in unmatched_ckpt_keys:333        result_state_dict[k] = ckpt_state_dict[k]334    return result_state_dict335 336 337def _group_keys_by_module(keys: List[str], original_names: Dict[str, str]):338    """339    Params in the same submodule are grouped together.340 341    Args:342        keys: names of all parameters343        original_names: mapping from parameter name to their name in the checkpoint344 345    Returns:346        dict[name -> all other names in the same group]347    """348 349    def _submodule_name(key):350        pos = key.rfind(".")351        if pos < 0:352            return None353        prefix = key[: pos + 1]354        return prefix355 356    all_submodules = [_submodule_name(k) for k in keys]357    all_submodules = [x for x in all_submodules if x]358    all_submodules = sorted(all_submodules, key=len)359 360    ret = {}361    for prefix in all_submodules:362        group = [k for k in keys if k.startswith(prefix)]363        if len(group) <= 1:364            continue365        original_name_lcp = _longest_common_prefix_str([original_names[k] for k in group])366        if len(original_name_lcp) == 0:367            # don't group weights if original names don't share prefix368            continue369 370        for k in group:371            if k in ret:372                continue373            ret[k] = group374    return ret375 376 377def _longest_common_prefix(names: List[str]) -> str:378    """379    ["abc.zfg", "abc.zef"] -> "abc."380    """381    names = [n.split(".") for n in names]382    m1, m2 = min(names), max(names)383    ret = [a for a, b in zip(m1, m2) if a == b]384    ret = ".".join(ret) + "." if len(ret) else ""385    return ret386 387 388def _longest_common_prefix_str(names: List[str]) -> str:389    m1, m2 = min(names), max(names)390    lcp = []391    for a, b in zip(m1, m2):392        if a == b:393            lcp.append(a)394        else:395            break396    lcp = "".join(lcp)397    return lcp398 399 400def _group_str(names: List[str]) -> str:401    """402    Turn "common1", "common2", "common3" into "common{1,2,3}"403    """404    lcp = _longest_common_prefix_str(names)405    rest = [x[len(lcp) :] for x in names]406    rest = "{" + ",".join(rest) + "}"407    ret = lcp + rest408 409    # add some simplification for BN specifically410    ret = ret.replace("bn_{beta,running_mean,running_var,gamma}", "bn_*")411    ret = ret.replace("bn_beta,bn_running_mean,bn_running_var,bn_gamma", "bn_*")412    return ret413