Initial commit
Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
commit
71db2d1ab9
126 changed files with 3198 additions and 0 deletions
82
saved/image_recognition_spike/match_spike_dinov2.py
Normal file
82
saved/image_recognition_spike/match_spike_dinov2.py
Normal file
|
|
@ -0,0 +1,82 @@
|
|||
"""
|
||||
Same spike as match_spike.py, but using DINOv2 instead of CLIP.
|
||||
|
||||
DINOv2 is trained with a self-supervised image-only objective specifically
|
||||
aimed at instance/fine-grained visual similarity, which is a better fit for
|
||||
"is this query photo the same object as this reference photo" than CLIP
|
||||
(CLIP is trained for text-image alignment and tends to cluster images by
|
||||
generic scene/style rather than object identity).
|
||||
|
||||
Usage: .venv\\Scripts\\python.exe match_spike_dinov2.py
|
||||
"""
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from PIL import Image
|
||||
from transformers import AutoImageProcessor, AutoModel
|
||||
|
||||
HERE = Path(__file__).parent
|
||||
REF_DIR = HERE / "images" / (sys.argv[1] if len(sys.argv) > 1 else "reference")
|
||||
QUERY_DIR = HERE / "images" / "query"
|
||||
|
||||
MODEL_NAME = "facebook/dinov2-base"
|
||||
|
||||
|
||||
def load_model():
|
||||
processor = AutoImageProcessor.from_pretrained(MODEL_NAME)
|
||||
model = AutoModel.from_pretrained(MODEL_NAME)
|
||||
model.eval()
|
||||
return model, processor
|
||||
|
||||
|
||||
def embed_image(model, processor, path: Path) -> torch.Tensor:
|
||||
image = Image.open(path).convert("RGB")
|
||||
inputs = processor(images=image, return_tensors="pt")
|
||||
with torch.no_grad():
|
||||
outputs = model(**inputs)
|
||||
features = outputs.last_hidden_state[:, 0, :] # CLS token
|
||||
features = features / features.norm(dim=-1, keepdim=True)
|
||||
return features.squeeze(0)
|
||||
|
||||
|
||||
def main():
|
||||
print(f"Loading {MODEL_NAME}...")
|
||||
model, processor = load_model()
|
||||
|
||||
ref_paths = (
|
||||
sorted(REF_DIR.glob("*.jpg"))
|
||||
+ sorted(REF_DIR.glob("*.webp"))
|
||||
+ sorted(REF_DIR.glob("*.png"))
|
||||
)
|
||||
query_paths = (
|
||||
sorted(QUERY_DIR.glob("*.jpg"))
|
||||
+ sorted(QUERY_DIR.glob("*.webp"))
|
||||
+ sorted(QUERY_DIR.glob("*.png"))
|
||||
)
|
||||
|
||||
print(f"Embedding {len(ref_paths)} reference images...")
|
||||
ref_names = [p.stem for p in ref_paths]
|
||||
ref_embeds = torch.stack([embed_image(model, processor, p) for p in ref_paths])
|
||||
|
||||
print(f"Embedding {len(query_paths)} query images...\n")
|
||||
correct = 0
|
||||
for qpath in query_paths:
|
||||
qembed = embed_image(model, processor, qpath)
|
||||
sims = ref_embeds @ qembed
|
||||
ranked = sorted(zip(ref_names, sims.tolist()), key=lambda x: -x[1])
|
||||
top_name, top_sim = ranked[0]
|
||||
expected = qpath.stem.split("_")[0]
|
||||
is_correct = top_name == expected
|
||||
correct += is_correct
|
||||
marker = "OK " if is_correct else "MISS"
|
||||
print(f"[{marker}] query={qpath.stem:<18} -> best={top_name:<10} sim={top_sim:.4f}")
|
||||
runner_up = ", ".join(f"{n}={s:.3f}" for n, s in ranked[1:4])
|
||||
print(f" runner-up: {runner_up}")
|
||||
|
||||
print(f"\n{correct}/{len(query_paths)} correct top-1 matches")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Loading…
Add table
Add a link
Reference in a new issue