Buckets:
| #!/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.