MCPcopy Create free account
hub / github.com/tensorpack/tensorpack / convert_weights

Function convert_weights

examples/FasterRCNN/convert_d2/convert_d2.py:33–122  ·  view source on GitHub ↗
(d, cfg)

Source from the content-addressed store, hash-verified

31
32
33def convert_weights(d, cfg):
34 has_fpn = "fpn" in cfg.MODEL.BACKBONE.NAME
35
36 ret = {}
37
38 def _convert_conv(src, dst):
39 src_w = d.pop(src + ".weight").transpose(2, 3, 1, 0)
40 ret[dst + "/W"] = src_w
41 if src + ".norm.weight" in d: # has norm
42 ret[dst + "/bn/gamma"] = d.pop(src + ".norm.weight")
43 ret[dst + "/bn/beta"] = d.pop(src + ".norm.bias")
44 ret[dst + "/bn/variance/EMA"] = d.pop(src + ".norm.running_var")
45 ret[dst + "/bn/mean/EMA"] = d.pop(src + ".norm.running_mean")
46 if src + ".bias" in d:
47 ret[dst + "/b"] = d.pop(src + ".bias")
48
49 def _convert_fc(src, dst):
50 ret[dst + "/W"] = d.pop(src + ".weight").transpose()
51 ret[dst + "/b"] = d.pop(src + ".bias")
52
53 if has_fpn:
54 backbone_prefix = "backbone.bottom_up."
55 else:
56 backbone_prefix = "backbone."
57 _convert_conv(backbone_prefix + "stem.conv1", "conv0")
58 for grpid in range(4):
59 if not has_fpn and grpid == 3:
60 backbone_prefix = "roi_heads."
61 for blkid in range([3, 4, 6 if cfg.MODEL.RESNETS.DEPTH == 50 else 23, 3][grpid]):
62 _convert_conv(backbone_prefix + f"res{grpid + 2}.{blkid}.conv1",
63 f"group{grpid}/block{blkid}/conv1")
64 _convert_conv(backbone_prefix + f"res{grpid + 2}.{blkid}.conv2",
65 f"group{grpid}/block{blkid}/conv2")
66 _convert_conv(backbone_prefix + f"res{grpid + 2}.{blkid}.conv3",
67 f"group{grpid}/block{blkid}/conv3")
68 if blkid == 0:
69 _convert_conv(backbone_prefix + f"res{grpid + 2}.{blkid}.shortcut",
70 f"group{grpid}/block{blkid}/convshortcut")
71
72 if has_fpn:
73 for lvl in range(2, 6):
74 _convert_conv(f"backbone.fpn_lateral{lvl}", f"fpn/lateral_1x1_c{lvl}")
75 _convert_conv(f"backbone.fpn_output{lvl}", f"fpn/posthoc_3x3_p{lvl}")
76
77 # RPN:
78 _convert_conv("proposal_generator.rpn_head.conv", "rpn/conv0")
79 _convert_conv("proposal_generator.rpn_head.objectness_logits", "rpn/class")
80 _convert_conv("proposal_generator.rpn_head.anchor_deltas", "rpn/box")
81
82 def _convert_box_predictor(src, dst):
83 if cfg.MODEL.ROI_BOX_HEAD.CLS_AGNOSTIC_BBOX_REG:
84 _convert_fc(src + ".bbox_pred", dst + "/box")
85 else:
86 v = d.pop(src + ".bbox_pred.bias")
87 ret[dst + "/box/b"] = np.concatenate((v[:4], v))
88 v = d.pop(src + ".bbox_pred.weight")
89 ret[dst + "/box/W"] = np.concatenate((v[:4, :], v), axis=0).transpose()
90

Callers 1

convert_d2.pyFile · 0.85

Calls 3

_convert_convFunction · 0.85
_convert_fcFunction · 0.85
_convert_box_predictorFunction · 0.85

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…