MCPcopy Create free account
hub / github.com/TPCD/DCCL / transform_moco_state_dict

Function transform_moco_state_dict

project_utils/general_utils.py:93–120  ·  view source on GitHub ↗

:param obj: Moco State Dict :param args: argsparse object with training classes :return: State dict compatable with standard resnet architecture

(obj, num_classes)

Source from the content-addressed store, hash-verified

91 return head_state_dict
92
93def transform_moco_state_dict(obj, num_classes):
94
95 """
96 :param obj: Moco State Dict
97 :param args: argsparse object with training classes
98 :return: State dict compatable with standard resnet architecture
99 """
100
101 newmodel = {}
102 for k, v in obj.items():
103 if not k.startswith("module.encoder_q."):
104 continue
105 old_k = k
106 k = k.replace("module.encoder_q.", "")
107
108 if k.startswith("fc.2"):
109 continue
110
111 if k.startswith("fc.0"):
112 k = k.replace("0.", "")
113 if "weight" in k:
114 v = torch.randn((num_classes, v.size(1)))
115 elif "bias" in k:
116 v = torch.randn((num_classes,))
117
118 newmodel[k] = v
119
120 return newmodel
121
122
123def init_experiment(args, runner_name=None, exp_id=None):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected