| """
|
| eval_aifnet.py — standalone evaluation for AIFNet / SpiceFusionNet checkpoints.
|
|
|
| Reports overall + per-source Top-1 and macro-F1 on a manifest split. With
|
| --synthetic-shift, applies a controlled "consumer-camera" acquisition transform
|
| (JPEG compression + warm color cast + gamma + downscale) to every test image —
|
| a held-out synthetic domain for OOD generalization evidence when a real 3rd
|
| source is unavailable. Compare a robust model (AIFNet) vs an ERM ablation:
|
| the smaller the clean->shifted drop, the more acquisition-invariant the model.
|
|
|
| python eval_aifnet.py --arch aifnet --ckpt outputs/checkpoints/aifnet_overlap_aif/best.pth
|
| python eval_aifnet.py --arch aifnet --ckpt .../best.pth --synthetic-shift
|
| """
|
| import argparse
|
| import json
|
| from collections import defaultdict
|
| from pathlib import Path
|
|
|
| import numpy as np
|
| from PIL import Image
|
| import torch
|
| import albumentations as A
|
| from albumentations.pytorch import ToTensorV2
|
| from sklearn.metrics import balanced_accuracy_score, f1_score
|
|
|
| import config
|
| from src.features import extract_all
|
| from src.dataset import load_manifest_splits, source_of
|
| from eval_source_probe import load_model
|
|
|
|
|
| def synthetic_shift():
|
| """Controlled 'consumer-camera' acquisition shift (albumentations 2.x API)."""
|
| return A.Compose([
|
| A.ImageCompression(quality_range=(25, 45), p=1.0),
|
| A.ColorJitter(brightness=0.25, contrast=0.25, saturation=0.35, hue=0.06, p=1.0),
|
| A.RandomGamma(gamma_limit=(70, 130), p=1.0),
|
| A.Downscale(scale_range=(0.5, 0.75), p=0.7),
|
| ])
|
|
|
|
|
| def normalize():
|
| s = int(config.IMG_SIZE * 256 / 224)
|
| return A.Compose([A.Resize(s, s), A.CenterCrop(config.IMG_SIZE, config.IMG_SIZE),
|
| A.Normalize(mean=config.IMG_MEAN, std=config.IMG_STD), ToTensorV2()])
|
|
|
|
|
| @torch.no_grad()
|
| def evaluate(model, samples, device, batch, shift_tf, norm_tf, arch):
|
| model.eval()
|
| y_true, y_pred, srcs = [], [], []
|
| img_b, tex_b, col_b, meta_b = [], [], [], []
|
|
|
| def flush():
|
| if not img_b:
|
| return
|
| x = torch.stack(img_b).to(device)
|
| tex = torch.from_numpy(np.stack(tex_b)).float().to(device)
|
| col = torch.from_numpy(np.stack(col_b)).float().to(device)
|
| logits = model(x, tex, col) if arch == "aifnet" else model.forward_fusion(x, tex, col)[0]
|
| pred = logits.argmax(1).cpu().numpy()
|
| for p, (lab, s) in zip(pred, meta_b):
|
| y_pred.append(int(p)); y_true.append(lab); srcs.append(s)
|
| img_b.clear(); tex_b.clear(); col_b.clear(); meta_b.clear()
|
|
|
| for path, label in samples:
|
| try:
|
| img = np.array(Image.open(path).convert("RGB"))
|
| except Exception:
|
| continue
|
| if shift_tf is not None:
|
| img = shift_tf(image=img)["image"]
|
| tex_np, col_np = extract_all(img)
|
| img_b.append(norm_tf(image=img)["image"])
|
| tex_b.append(tex_np); col_b.append(col_np)
|
| meta_b.append((int(label), source_of(path)))
|
| if len(img_b) >= batch:
|
| flush()
|
| flush()
|
|
|
| y_true, y_pred, srcs = np.array(y_true), np.array(y_pred), np.array(srcs)
|
| overall = float((y_true == y_pred).mean())
|
| f1m = float(f1_score(y_true, y_pred, average="macro"))
|
| per_src = {s: float((y_true[srcs == s] == y_pred[srcs == s]).mean())
|
| for s in sorted(set(srcs))}
|
| gap = (max(per_src.values()) - min(per_src.values())) if len(per_src) > 1 else 0.0
|
| return {"top1": overall, "macro_f1": f1m, "per_source_top1": per_src, "source_gap": gap,
|
| "n": int(len(y_true))}
|
|
|
|
|
| def main():
|
| ap = argparse.ArgumentParser()
|
| ap.add_argument("--ckpt", required=True)
|
| ap.add_argument("--arch", default="aifnet", choices=["aifnet", "spicefusion"])
|
| ap.add_argument("--manifest", default="outputs/manifest_overlap_only.json")
|
| ap.add_argument("--split", default="test", choices=["train", "val", "test"])
|
| ap.add_argument("--synthetic-shift", action="store_true")
|
| ap.add_argument("--batch", type=int, default=32)
|
| ap.add_argument("--out", default=None)
|
| ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
|
| args = ap.parse_args()
|
|
|
| device = torch.device(args.device)
|
| splits, classes = load_manifest_splits(args.manifest)
|
| samples = list(zip(*splits[args.split]))
|
| samples = [(p, y) for p, y in samples]
|
| model = load_model(args.ckpt, device, args.arch)
|
| norm_tf = normalize()
|
|
|
| print(f"eval {Path(args.ckpt).name} [{args.arch}] on {args.manifest} {args.split} "
|
| f"({len(samples)} samples)")
|
| clean = evaluate(model, samples, device, args.batch, None, norm_tf, args.arch)
|
| print(f" CLEAN top1={clean['top1']:.4f} macroF1={clean['macro_f1']:.4f} "
|
| f"gap={clean['source_gap']:.4f} per_src={ {k:round(v,3) for k,v in clean['per_source_top1'].items()} }")
|
| out = {"checkpoint": args.ckpt, "arch": args.arch, "clean": clean}
|
|
|
| if args.synthetic_shift:
|
| shifted = evaluate(model, samples, device, args.batch, synthetic_shift(), norm_tf, args.arch)
|
| drop = clean["top1"] - shifted["top1"]
|
| print(f" SHIFTED top1={shifted['top1']:.4f} macroF1={shifted['macro_f1']:.4f} "
|
| f"gap={shifted['source_gap']:.4f} | clean->shift drop = {drop:.4f}")
|
| out["synthetic_shift"] = shifted
|
| out["shift_drop_top1"] = drop
|
|
|
| if args.out:
|
| Path(args.out).parent.mkdir(parents=True, exist_ok=True)
|
| json.dump(out, open(args.out, "w"), indent=2)
|
| print(f" -> {args.out}")
|
|
|
|
|
| if __name__ == "__main__":
|
| main()
|
|
|