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