83 lines
2.7 KiB
Python
83 lines
2.7 KiB
Python
"""
|
|
De-risking spike: can a generic vision embedding model (CLIP) tell Pokemon
|
|
figures/toys apart via nearest-neighbor lookup, with no training?
|
|
|
|
Pipeline: embed every image under images/reference/ (one file per Pokemon,
|
|
filename = class name) to build a small reference index, then embed every
|
|
image under images/query/ and report the nearest reference neighbor +
|
|
cosine similarity. A "pass" is query/<name>.jpg matching reference/<name>.jpg
|
|
as the top-1 hit.
|
|
|
|
Usage: .venv\\Scripts\\python.exe match_spike.py
|
|
"""
|
|
|
|
import sys
|
|
from pathlib import Path
|
|
|
|
import open_clip
|
|
import torch
|
|
from PIL import Image
|
|
|
|
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 = "ViT-B-32-quickgelu"
|
|
PRETRAINED = "openai"
|
|
|
|
|
|
def load_model():
|
|
model, _, preprocess = open_clip.create_model_and_transforms(
|
|
MODEL_NAME, pretrained=PRETRAINED
|
|
)
|
|
model.eval()
|
|
return model, preprocess
|
|
|
|
|
|
def embed_image(model, preprocess, path: Path) -> torch.Tensor:
|
|
image = preprocess(Image.open(path).convert("RGB")).unsqueeze(0)
|
|
with torch.no_grad():
|
|
features = model.encode_image(image)
|
|
features = features / features.norm(dim=-1, keepdim=True)
|
|
return features.squeeze(0)
|
|
|
|
|
|
def main():
|
|
print(f"Loading {MODEL_NAME} ({PRETRAINED})...")
|
|
model, preprocess = 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, preprocess, 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, preprocess, qpath)
|
|
sims = ref_embeds @ qembed # cosine similarity, both sides unit-norm
|
|
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()
|