#!/usr/bin/env python3
"""Test DINOv2 su FOTO REALI dei Santuari (non aumentate).

- Galleria = 8 santuari distinti (foto ritagliate).
- Matrice di similarita (carte diverse devono stare BASSE; focus sul trio
  "scogliera" 6583/6585/6586 che a occhio si somigliano).
- Match rotation-invariant: ogni carta in 4 rotazioni + jitter luce -> ritrova se stessa?
- Prova gemella: 6587 (stessa carta di 6588, sfondo+inquadratura diversi) -> trova 6588?
"""
import os, glob, itertools, numpy as np
from PIL import Image, ImageEnhance
import torch
from transformers import AutoImageProcessor, AutoModel

CROPS = "/tmp/fw_crops"
# 8 carte distinte (6587 escluso: e la gemella di 6588, la useremo come query)
GALLERY_IDS = ["IMG_6580", "IMG_6581", "IMG_6582", "IMG_6583",
               "IMG_6584", "IMG_6585", "IMG_6586", "IMG_6588"]
LABEL = {
    "IMG_6580": "Deserto/okiko", "IMG_6581": "Foresta/uddu",
    "IMG_6582": "Fiume/okiko",   "IMG_6583": "Rifugi 1xblu|verde",
    "IMG_6584": "Fiume -1xuddu", "IMG_6585": "Rifugi 2xclue",
    "IMG_6586": "Rifugi okiko+clue", "IMG_6588": "Rifugi uddu+goldlog",
}

MODEL = "facebook/dinov2-small"
print("Carico", MODEL, "...")
proc = AutoImageProcessor.from_pretrained(MODEL)
model = AutoModel.from_pretrained(MODEL).eval()


@torch.no_grad()
def embed(img):
    inp = proc(images=img.convert("RGB"), return_tensors="pt")
    v = model(**inp).last_hidden_state[:, 0]
    v = torch.nn.functional.normalize(v, dim=-1)
    return v[0].cpu().numpy()


def rots(img):
    return [img, img.rotate(90, expand=True), img.rotate(180, expand=True),
            img.rotate(270, expand=True)]


def embed_rotinv(path):
    """Lista di embedding per le 4 rotazioni (match = max su queste)."""
    im = Image.open(path)
    return [embed(r) for r in rots(im)]


def sim_rotinv(emb_list_a, emb_list_b):
    return max(float(np.dot(a, b)) for a in emb_list_a for b in emb_list_b)


# --- galleria (rotation-invariant) ---
gal = {i: embed_rotinv(os.path.join(CROPS, i + ".jpg")) for i in GALLERY_IDS}

print("\n=== Matrice similarita (rotation-invariant) — diagonale=1, fuori basso ===")
hdr = "             " + " ".join(f"{i[-4:]:>5}" for i in GALLERY_IDS)
print(hdr)
worst_off = 0.0
for a in GALLERY_IDS:
    cells = []
    for b in GALLERY_IDS:
        s = sim_rotinv(gal[a], gal[b])
        cells.append(f"{s:.2f}")
        if a != b:
            worst_off = max(worst_off, s)
    print(f"{a[-4:]:>4} {LABEL[a][:8]:<8} " + " ".join(f"{c:>5}" for c in cells))
print(f"\nMax similarita tra carte DIVERSE: {worst_off:.3f}  (piu e bassa, meglio e)")

# margine peggiore: per ogni carta, gap tra se stessa e il vicino piu vicino
print("\n=== Margine self vs 2° classificato ===")
min_margin = 9.0
for a in GALLERY_IDS:
    others = sorted(((sim_rotinv(gal[a], gal[b]), b) for b in GALLERY_IDS if b != a), reverse=True)
    nn_sim, nn = others[0]
    margin = 1.0 - nn_sim
    min_margin = min(min_margin, margin)
    flag = "  <-- stretto" if margin < 0.15 else ""
    print(f"  {a[-4:]} {LABEL[a]:<18} vicino+={nn[-4:]}({LABEL[nn]}) sim={nn_sim:.2f} margine={margin:.2f}{flag}")
print(f"Margine minimo: {min_margin:.3f}")

# --- test match con jitter (luce/contrasto) + rotazioni casuali ---
print("\n=== Match su versioni 'disturbate' (luce+contrasto+rotazione) ===")
import random
tot = ok = 0
for a in GALLERY_IDS:
    base = Image.open(os.path.join(CROPS, a + ".jpg")).convert("RGB")
    for k in range(4):
        random.seed(hash(a) % 1000 + k)
        im = base.rotate(random.choice([0, 90, 180, 270]), expand=True)
        im = ImageEnhance.Brightness(im).enhance(random.uniform(0.7, 1.3))
        im = ImageEnhance.Contrast(im).enhance(random.uniform(0.8, 1.2))
        q = [embed(im)]
        scored = sorted(((sim_rotinv(q, gal[b]), b) for b in GALLERY_IDS), reverse=True)
        hit = scored[0][1] == a
        tot += 1; ok += hit
        if not hit:
            print(f"  ❌ {a[-4:]} k{k}: 1°={scored[0][1][-4:]}({scored[0][0]:.2f}) "
                  f"2°={scored[1][1][-4:]}({scored[1][0]:.2f})")
print(f"Accuratezza: {ok}/{tot} = {100*ok/tot:.0f}%")

# --- prova gemella reale: 6587 (sfondo beige, vista parziale) -> 6588? ---
print("\n=== Gemella reale: 6587 (beige, ritaglio parziale) deve trovare 6588 ===")
q587 = embed_rotinv(os.path.join(CROPS, "IMG_6587.jpg"))
scored = sorted(((sim_rotinv(q587, gal[b]), b) for b in GALLERY_IDS), reverse=True)
for s, b in scored[:4]:
    mark = " ✅" if b == "IMG_6588" else ""
    print(f"  {b[-4:]} {LABEL[b]:<18} {s:.3f}{mark}")
print(f"\nRisultato: 6587 -> {scored[0][1][-4:]} ({'CORRETTO' if scored[0][1]=='IMG_6588' else 'SBAGLIATO'})")
