1import os
2import hashlib
3
4import torch
5import timm
6
7
8# download weights
9url = "https://huggingface.co/eplekh/secoeco/resolve/main/ablation_B12_weights.ckpt"
10ckpt = torch.hub.load_state_dict_from_url(url, map_location="cpu", progress=True)
11arch, image_size, bands = ckpt["hyper_parameters"]["arch"], ckpt["hyper_parameters"]["in_size"], ckpt["hyper_parameters"]["bands"]
12print(arch, image_size, bands)
13
14# Bands correspond to B9 from https://github.com/PlekhanovaElena/ssl4eco/blob/7445e048035f7ae31c0eb45e1ed8426c9989fe56/pretraining/pretrain_seco_3heads.py#L220
15bands = ['B1', 'B2', 'B3', 'B4', 'B5', 'B6', 'B7', 'B8', 'B8A', 'B9', 'B11', 'B12']
16
17# map weights to timm and torchvision compatible
18layer_mapping = {
19 "0" : "conv1",
20 "1" : "bn1",
21 "4" : "layer1",
22 "5" : "layer2",
23 "6" : "layer3",
24 "7" : "layer4",
25}
26state_dict = {k.replace("encoder_q.", ""): v for k, v in ckpt["state_dict"].items() if k.startswith("encoder_q.")}
27state_dict = {k.replace(k.split(".")[0], layer_mapping[k.split(".")[0]]): v for k, v in state_dict.items()}
28
29model = timm.create_model("resnet50", pretrained=False, in_chans=len(bands), num_classes=0)
30model.load_state_dict(state_dict, strict=True)
31
32# save and compute hash
33filename = "resnet50_sentinel2_all_seco_eco.pth"
34torch.save(model.state_dict(), filename)
35md5 = hashlib.md5(open(filename, "rb").read()).hexdigest()[:8]
36os.rename(filename, filename.replace(".pth", f"-{md5}.pth"))