Sor0ush's picture
download
raw
8.09 kB
#!/usr/bin/env python3
"""Evaluate encoders on binding / same-diff / geometry / JER."""
from __future__ import annotations
import argparse
import json
import sys
import time
import traceback
from pathlib import Path
import numpy as np
import torch
from PIL import Image
ROOT = Path(__file__).resolve().parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
from binding_task import ( # noqa: E402
generate_binding_trials,
generate_samediff_trials,
verify_disjoint_colors,
)
from metrics import ( # noqa: E402
evaluate_binding,
evaluate_samediff,
global_geometry,
jacobian_effective_rank_noise,
local_isotropy,
)
from models import load_encoder, suite_keys # noqa: E402
def load_geometry_images(n: int, seed: int = 42) -> list[Image.Image]:
"""Load natural images for geometry metrics (public HF sets; paper uses ImageNet val)."""
from datasets import load_dataset
candidates = [
("uoft-cs/cifar100", "test", "img"),
("zh-plus/tiny-imagenet", "valid", "image"),
]
rng = np.random.default_rng(seed)
for repo, split, col in candidates:
try:
ds = load_dataset(repo, split=split)
idxs = rng.choice(len(ds), size=min(n, len(ds)), replace=False)
images = []
for i in idxs:
im = ds[int(i)][col]
if not isinstance(im, Image.Image):
im = Image.fromarray(np.asarray(im))
images.append(im.convert("RGB"))
print(f"[geometry] loaded {len(images)} images from {repo}:{split}", flush=True)
return images
except Exception as e:
print(f"[geometry] {repo} failed ({e})", flush=True)
print("[geometry] falling back to procedural textures", flush=True)
images = []
for _ in range(n):
arr = rng.integers(0, 255, size=(224, 224, 3), dtype=np.uint8)
images.append(Image.fromarray(arr, mode="RGB"))
return images
def embed_many(encoder, images: list[Image.Image], batch_size: int = 32) -> np.ndarray:
outs = []
for i in range(0, len(images), batch_size):
outs.append(encoder.encode(images[i : i + batch_size]))
return np.concatenate(outs, axis=0)
def evaluate_one(
key: str,
*,
device: str,
n_binding: int,
n_samediff: int,
n_geom: int,
n_jer: int,
jer_k: int,
skip_jer: bool,
skip_geom: bool,
) -> dict:
t0 = time.time()
enc = load_encoder(key, device=device)
print(f"[load] {enc.name} on {enc.device}", flush=True)
binding_trials = generate_binding_trials(n_binding, seed=42)
samediff_trials = generate_samediff_trials(n_samediff, seed=42)
disjoint_rate = verify_disjoint_colors(binding_trials)
binding_acc = evaluate_binding(enc.encode, binding_trials)
disc_acc = evaluate_samediff(enc.encode, samediff_trials)
print(f"[task] binding={binding_acc:.4f} disc={disc_acc:.4f} disjoint={disjoint_rate:.3f}", flush=True)
row = {
"key": key,
"model": enc.name,
"family": enc.family,
"binding": binding_acc,
"disc": disc_acc,
"disjoint_color_rate": disjoint_rate,
"G.PR": None,
"G.Iso": None,
"L.Iso": None,
"JER": None,
"seconds": None,
"error": None,
}
if not skip_geom:
geom_images = load_geometry_images(n_geom, seed=42)
embs = embed_many(enc, geom_images)
g = global_geometry(embs)
liso = local_isotropy(embs, k=32, n_anchors=min(500, len(embs)), seed=42)
row["G.PR"] = g["G.PR"]
row["G.Iso"] = g["G.Iso"]
row["L.Iso"] = liso
print(f"[geom] G.PR={g['G.PR']:.4f} G.Iso={g['G.Iso']:.4f} L.Iso={liso:.4f}", flush=True)
if not skip_jer:
try:
jer = jacobian_effective_rank_noise(
enc.encode_tensor,
n_images=n_jer,
k=jer_k,
device=enc.device,
seed=42,
)
row["JER"] = jer
print(f"[jer] JER={jer:.4f}", flush=True)
except Exception as e:
row["error"] = f"JER failed: {type(e).__name__}: {e}"
print(f"[jer] FAILED: {row['error']}", flush=True)
row["seconds"] = time.time() - t0
return row
def main() -> None:
p = argparse.ArgumentParser()
p.add_argument("--suite", default="core", choices=["smoke", "claim4", "core", "full", "custom"])
p.add_argument("--models", nargs="*", default=None)
p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
p.add_argument("--n-binding", type=int, default=500)
p.add_argument("--n-samediff", type=int, default=500)
p.add_argument("--n-geom", type=int, default=1000)
p.add_argument("--n-jer", type=int, default=100)
p.add_argument("--jer-k", type=int, default=32)
p.add_argument("--skip-jer", action="store_true")
p.add_argument("--skip-geom", action="store_true")
p.add_argument("--out", type=Path, default=ROOT / "outputs" / "results.json")
args = p.parse_args()
keys = args.models if args.suite == "custom" or args.models else suite_keys(args.suite)
args.out.parent.mkdir(parents=True, exist_ok=True)
rows = []
if args.out.exists():
try:
prev = json.loads(args.out.read_text())
rows = prev.get("rows", [])
# Retry keys that failed or lack binding; keep successful rows.
done = {
r["key"]
for r in rows
if r.get("binding") is not None and (r.get("error") is None or str(r.get("error", "")).startswith("JER failed"))
and (args.skip_jer or r.get("JER") is not None or str(r.get("error", "")).startswith("JER failed"))
}
# Simpler: done only if binding present and (JER present or skip_jer or non-JER error absent)
done = set()
for r in rows:
if r.get("binding") is None:
continue
if args.skip_jer or r.get("JER") is not None:
done.add(r["key"])
keys = [k for k in keys if k not in done]
print(f"[resume] {len(done)} done, {len(keys)} remaining", flush=True)
except Exception:
rows = []
meta = {
"paper": "2602.03282",
"openreview": "gMwcZb18s7",
"suite": args.suite,
"device": args.device,
"n_binding": args.n_binding,
"n_samediff": args.n_samediff,
"n_geom": args.n_geom,
"n_jer": args.n_jer,
"jer_k": args.jer_k,
"cuda": torch.cuda.is_available(),
"gpu": torch.cuda.get_device_name(0) if torch.cuda.is_available() else None,
}
print(json.dumps(meta, indent=2), flush=True)
for key in keys:
try:
row = evaluate_one(
key,
device=args.device,
n_binding=args.n_binding,
n_samediff=args.n_samediff,
n_geom=args.n_geom,
n_jer=args.n_jer,
jer_k=args.jer_k,
skip_jer=args.skip_jer,
skip_geom=args.skip_geom,
)
except Exception as e:
traceback.print_exc()
row = {
"key": key,
"model": key,
"family": None,
"binding": None,
"disc": None,
"disjoint_color_rate": None,
"G.PR": None,
"G.Iso": None,
"L.Iso": None,
"JER": None,
"seconds": None,
"error": f"{type(e).__name__}: {e}",
}
rows = [r for r in rows if r.get("key") != key] + [row]
payload = {"meta": meta, "rows": rows}
args.out.write_text(json.dumps(payload, indent=2))
print(f"[saved] {args.out} ({len(rows)} rows)", flush=True)
print(json.dumps({"meta": meta, "rows": rows}, indent=2), flush=True)
if __name__ == "__main__":
main()

Xet Storage Details

Size:
8.09 kB
·
Xet hash:
a3bb6867b320d3164d5c52bc86c1ad43328d441094c528eb051a6495ba8496b8

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.