Views
No views yet
1
2import torch
3
4import torch.nn.functional as F
5
6from torchvision import transforms
7
8from PIL import Image
9
10import argparse
11
12import os
13
14
15
16# --------------------------------------------------
17
18# Configuration
19
20# --------------------------------------------------
21
22IMAGE_SIZE = (128, 128)
23
24MEAN = [0.485, 0.456, 0.406]
25
26STD = [0.229, 0.224, 0.225]
27
28
29
30
31
32# --------------------------------------------------
33
34# Image preprocessing (must match training!)
35
36# --------------------------------------------------
37
38transform = transforms.Compose([
39
40 transforms.Resize(IMAGE_SIZE),
41
42 transforms.ToTensor(),
43
44 transforms.Normalize(mean=MEAN, std=STD),
45
46])
47
48
49
50
51
52# --------------------------------------------------
53
54# Load ReID model
55
56# --------------------------------------------------
57
58def load_model(model_path, device):
59
60 """
61
62 Loads a trained Omni-Scale ReID model (.pt)
63
64 """
65
66 model = torch.load(model_path, map_location=device)
67
68
69
70 model.to(device)
71
72 model.eval()
73
74 return model
75
76
77
78
79
80# --------------------------------------------------
81
82# Extract embedding
83
84# --------------------------------------------------
85
86def extract_embedding(model, image_path, device):
87
88 """
89
90 Extracts a normalized ReID embedding from an image
91
92 """
93
94 if not os.path.exists(image_path):
95
96 raise FileNotFoundError(image_path)
97
98
99
100 img = Image.open(image_path).convert("RGB")
101
102 img = transform(img)
103
104 img = img.unsqueeze(0).to(device) # [1, 3, 128, 128]
105
106
107
108 with torch.no_grad():
109
110 feat = model(img)
111
112
113
114 # Flatten and L2-normalize (standard ReID practice)
115
116 feat = feat.squeeze(0)
117
118 feat = F.normalize(feat, dim=0)
119
120
121
122 return feat.cpu()
123
124
125
126
127
128# --------------------------------------------------
129
130# Cosine similarity
131
132# --------------------------------------------------
133
134def cosine_sim(feat1, feat2):
135
136 return F.cosine_similarity(
137
138 feat1.unsqueeze(0),
139
140 feat2.unsqueeze(0)
141
142 ).item()
143
144
145
146
147
148# --------------------------------------------------
149
150# Main
151
152# --------------------------------------------------
153
154def main():
155
156 parser = argparse.ArgumentParser("Omni-Scale ReID Inference")
157
158 parser.add_argument("--model", type=str, required=True, help="Path to .pt model")
159
160 parser.add_argument("--img1", type=str, required=True, help="Query image")
161
162 parser.add_argument("--img2", type=str, default=None, help="Gallery image (optional)")
163
164 args = parser.parse_args()
165
166
167
168 device = "cuda" if torch.cuda.is_available() else "cpu"
169
170 print(f"Using device: {device}")
171
172
173
174 # Load model
175
176 model = load_model(args.model, device)
177
178
179
180 # Extract embedding(s)
181
182 feat1 = extract_embedding(model, args.img1, device)
183
184 print(f"Embedding shape: {feat1.shape}")
185
186
187
188 if args.img2:
189
190 feat2 = extract_embedding(model, args.img2, device)
191
192 sim = cosine_sim(feat1, feat2)
193
194 print(f"Cosine similarity: {sim:.4f}")
195
196
197
198
199
200if __name__ == "__main__":
201
202 main()
203