forked from TimJaspers0801/SurgeNet
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathload_vit_models.py
More file actions
124 lines (105 loc) · 4.35 KB
/
Copy pathload_vit_models.py
File metadata and controls
124 lines (105 loc) · 4.35 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
import torch
import timm
# This code loads all DINO model weights (v1, v2, v3) using timm library. Some unexpected keys are ignored during loading, this is intended behaviour.
# Optionally specify classes if needed for classifcation tasks
n_classes = 10
# URLs for pre-trained DINO weights
urls = {
'path_dinov1_vits': 'https://huggingface.co/rlpddejong/SurgeNetXL_DINOv1-v3/resolve/main/DINOv1_ViTs16_size224_SurgeNetXL.pth?download=true',
'path_dinov1_vitb': 'https://huggingface.co/rlpddejong/SurgeNetXL_DINOv1-v3/resolve/main/DINOv1_ViTb16_size224_SurgeNetXL.pth?download=true',
'path_dinov2_vits': 'https://huggingface.co/rlpddejong/SurgeNetXL_DINOv1-v3/resolve/main/DINOv2_ViTs14_size336_SurgeNetXL.pth?download=true',
'path_dinov2_vitb': 'https://huggingface.co/rlpddejong/SurgeNetXL_DINOv1-v3/resolve/main/DINOv2_ViTb14_size336_SurgeNetXL.pth?download=true',
'path_dinov2_vitl': 'https://huggingface.co/rlpddejong/SurgeNetXL_DINOv1-v3/resolve/main/DINOv2_ViTl14_size336_SurgeNetXL.pth?download=true',
'path_dinov3_vits': 'https://huggingface.co/rlpddejong/SurgeNetXL_DINOv1-v3/resolve/main/DINOv3_ViTs16_size336_SurgeNetXL.pth?download=true',
'path_dinov3_vitb': 'https://huggingface.co/rlpddejong/SurgeNetXL_DINOv1-v3/resolve/main/DINOv3_ViTb16_size336_SurgeNetXL.pth?download=true',
'path_dinov3_vitl': 'https://huggingface.co/rlpddejong/SurgeNetXL_DINOv1-v3/resolve/main/DINOv3_ViTl16_size336_SurgeNetXL.pth?download=true',
}
###################################
### Loading dinov1(using timm) ###
###################################
# ViT-s
model = timm.create_model(
'vit_small_patch16_224.dino',
img_size=(224, 224),
patch_size=16,
num_classes=n_classes,
)
state_dict = torch.hub.load_state_dict_from_url(urls['path_dinov1_vits'])
msg = model.load_state_dict(state_dict, strict=False)
print("\nLoaded DINOv1 ViT-s weights with msg:\n", msg)
# ViT-b
model = timm.create_model(
'vit_base_patch16_224.dino',
img_size=(224, 224),
patch_size=16,
num_classes=n_classes,
)
state_dict = torch.hub.load_state_dict_from_url(urls['path_dinov1_vitb'])
msg = model.load_state_dict(state_dict, strict=False)
print("\nLoaded DINOv1 ViT-b weights with msg:\n", msg)
###################################
### Loading dinov2 (using timm) ###
###################################
# ViT-s
model = timm.create_model(
'vit_small_patch14_dinov2',
img_size=(336, 336),
patch_size=14,
num_classes=n_classes,
)
state_dict = torch.hub.load_state_dict_from_url(urls['path_dinov2_vits'])
msg = model.load_state_dict(state_dict, strict=False)
print("\nLoaded DINOv2 ViT-s weights with msg:\n", msg)
# ViT-b
model = timm.create_model(
'vit_base_patch14_dinov2',
img_size=(336, 336),
patch_size=14,
num_classes=n_classes,
)
state_dict = torch.hub.load_state_dict_from_url(urls['path_dinov2_vitb'])
msg = model.load_state_dict(state_dict, strict=False)
print("\nLoaded DINOv2 ViT-b weights with msg:\n", msg)
# ViT-l
model = timm.create_model(
'vit_large_patch14_dinov2',
img_size=(336, 336),
patch_size=14,
num_classes=n_classes,
)
state_dict = torch.hub.load_state_dict_from_url(urls['path_dinov2_vitl'])
msg = model.load_state_dict(state_dict, strict=False)
print("\nLoaded DINOv2 ViT-l weights with msg:\n", msg)
###########################################
### Loading dinov3 (using transformers) ###
###########################################
# ViTs
model = timm.create_model(
'vit_small_patch16_dinov3.lvd1689m',
img_size=(336, 336),
patch_size=16,
num_classes=n_classes,
)
state_dict = torch.hub.load_state_dict_from_url(urls['path_dinov3_vits'])
msg = model.load_state_dict(state_dict, strict=False)
print("\nLoaded DINOv3 ViT-s weights with msg:\n", msg)
# ViTb
model = timm.create_model(
'vit_base_patch16_dinov3.lvd1689m',
img_size=(336, 336),
patch_size=16,
num_classes=n_classes,
)
state_dict = torch.hub.load_state_dict_from_url(urls['path_dinov3_vitb'])
msg = model.load_state_dict(state_dict, strict=False)
print("\nLoaded DINOv3 ViT-b weights with msg:\n", msg)
# ViTl
model = timm.create_model(
'vit_large_patch16_dinov3.lvd1689m',
img_size=(336, 336),
patch_size=16,
num_classes=n_classes,
)
state_dict = torch.hub.load_state_dict_from_url(urls['path_dinov3_vitl'])
msg = model.load_state_dict(state_dict, strict=False)
print("\nLoaded DINOv3 ViT-l weights with msg:\n", msg)