:param obj: Moco State Dict :param args: argsparse object with training classes :return: State dict compatable with standard resnet architecture
(obj, num_classes)
| 91 | return head_state_dict |
| 92 | |
| 93 | def 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 | |
| 123 | def init_experiment(args, runner_name=None, exp_id=None): |
nothing calls this directly
no outgoing calls
no test coverage detected