(image_dir, image_dir_ref, device)
| 199 | |
| 200 | |
| 201 | def dinoeval_image(image_dir, image_dir_ref, device): |
| 202 | image_paths = [os.path.join(image_dir, path) for path in os.listdir(image_dir) |
| 203 | if path.endswith(('.png', '.jpg', '.jpeg', '.tiff', '.JPG'))] |
| 204 | image_paths_ref = [os.path.join(image_dir_ref, path) for path in os.listdir(image_dir_ref) |
| 205 | if path.endswith(('.png', '.jpg', '.jpeg', '.tiff', '.JPG'))] |
| 206 | |
| 207 | model = torch.hub.load('facebookresearch/dino:main', 'dino_vits16').to(device) |
| 208 | model.eval() |
| 209 | |
| 210 | image_feats = extract_all_images( |
| 211 | image_paths, model, DINOImageDataset, device, batch_size=64, num_workers=8) |
| 212 | |
| 213 | image_feats_ref = extract_all_images( |
| 214 | image_paths_ref, model, DINOImageDataset, device, batch_size=64, num_workers=8) |
| 215 | |
| 216 | image_feats = image_feats / \ |
| 217 | np.sqrt(np.sum(image_feats ** 2, axis=1, keepdims=True)) |
| 218 | image_feats_ref = image_feats_ref / \ |
| 219 | np.sqrt(np.sum(image_feats_ref ** 2, axis=1, keepdims=True)) |
| 220 | res = image_feats @ image_feats_ref.T |
| 221 | return np.mean(res) |
| 222 | |
| 223 | |
| 224 | def calmetrics(sample_root, target_paths, numgen, outpkl): |
no test coverage detected