Views
No views yet
miewid-msv3-latonia-1232.pt1import torch
2from PIL import Image
3from torchvision import transforms
4
5from transformers import AutoModel
6
7
8class ZoomCenterCrop:
9 def __init__(self, zoom=1.0):
10 self.zoom = zoom
11
12 def __call__(self, img):
13 w, h = img.size
14 m = int(min(h, w) / self.zoom)
15 left = (w - m) // 2
16 top = (h - m) // 2
17 return img.crop((left, top, left + m, top + m))
18
19
20preprocess = transforms.Compose([
21 ZoomCenterCrop(zoom=2.0),
22 transforms.Resize((440, 440)),
23 transforms.ToTensor(),
24 transforms.Normalize(
25 mean=[0.485, 0.456, 0.406],
26 std=[0.229, 0.224, 0.225]),
27])
28
29model = AutoModel.from_pretrained("conservationxlabs/miewid-msv3", trust_remote_code=True)
30ckpt = torch.load("miewid-msv3-latonia-1232.pt", map_location="cpu")
31model.load_state_dict(ckpt["model"], strict=True)
32model.eval()
33
34def embed(path):
35 img = Image.open(path).convert("RGB")
36 x = preprocess(img).unsqueeze(0)
37 with torch.no_grad():
38 emb = model(x)
39 return emb / emb.norm(dim=1, keepdim=True)
40
41e1 = embed("img1.jpg")
42e2 = embed("img2.jpg")
43cosine_sim = (e1 @ e2.T).item()
44print(cosine_sim)