""" evaluate.py — Evaluate trained models on the BDAPPV test split. Usage: # Segmentation python evaluate.py seg --provider google --checkpoint weights/deeplab_google_best.pth python evaluate.py seg --provider ign --checkpoint weights/deeplab_ign_best.pth # Classification python evaluate.py clf --provider google --checkpoint weights/inception_google_best.pth # Cross-provider (train Google, eval IGN) — the distribution shift benchmark python evaluate.py seg --provider ign --checkpoint weights/deeplab_google_best.pth """ import argparse import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision.models.segmentation import deeplabv3_resnet101 from torchvision.models import inception_v3 from datasets import load_dataset from dataset import SegmentationDataset, ClassificationDataset # ── Metrics ─────────────────────────────────────────────────────────────────── def compute_seg_metrics(pred_logits, target, threshold=0.5): pred = (torch.sigmoid(pred_logits) > threshold).float() tp = (pred * target).sum().item() fp = (pred * (1 - target)).sum().item() fn = ((1 - pred) * target).sum().item() iou = tp / max(tp + fp + fn, 1e-6) prec = tp / max(tp + fp, 1e-6) rec = tp / max(tp + fn, 1e-6) f1 = 2 * prec * rec / max(prec + rec, 1e-6) return iou, f1 def compute_clf_metrics(logits, labels, threshold=0.0): preds = (logits > threshold).float() acc = (preds == labels).float().mean().item() tp = (preds * labels).sum().item() fp = (preds * (1 - labels)).sum().item() fn = ((1 - preds) * labels).sum().item() prec = tp / max(tp + fp, 1e-6) rec = tp / max(tp + fn, 1e-6) f1 = 2 * prec * rec / max(prec + rec, 1e-6) return acc, prec, rec, f1 # ── Segmentation eval ───────────────────────────────────────────────────────── def eval_seg(args): device = ( torch.device("mps") if torch.backends.mps.is_available() else torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu") ) ds = load_dataset("gabrielkasmi/bdappv", args.provider) test_ds = SegmentationDataset(ds["test"], img_size=args.img_size, augment=False) loader = DataLoader(test_ds, batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers) model = deeplabv3_resnet101(weights=None, aux_loss=False) model.classifier[-1] = nn.Conv2d(256, 1, kernel_size=1) state = torch.load(args.checkpoint, map_location="cpu", weights_only=False) model_dict = model.state_dict() compatible = {k: v for k, v in state.items() if k in model_dict and v.shape == model_dict[k].shape} model_dict.update(compatible) model.load_state_dict(model_dict) model = model.to(device) model.eval() total_iou, total_f1 = 0.0, 0.0 with torch.no_grad(): for images, masks in loader: images, masks = images.to(device), masks.to(device) logits = model(images)["out"] if logits.shape[-2:] != masks.shape[-2:]: logits = nn.functional.interpolate(logits, size=masks.shape[-2:], mode="bilinear", align_corners=False) iou, f1 = compute_seg_metrics(logits, masks) total_iou += iou total_f1 += f1 n = len(loader) print(f"\n── Segmentation results ──────────────────────") print(f" Provider : {args.provider}") print(f" Checkpoint : {args.checkpoint}") print(f" Test images: {len(test_ds):,}") print(f" IoU : {total_iou / n:.4f}") print(f" F1 : {total_f1 / n:.4f}") # ── Classification eval ─────────────────────────────────────────────────────── def eval_clf(args): device = ( torch.device("mps") if torch.backends.mps.is_available() else torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu") ) ds = load_dataset("gabrielkasmi/bdappv", args.provider) test_ds = ClassificationDataset(ds["test"], img_size=299, augment=False) loader = DataLoader(test_ds, batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers) model = inception_v3(weights=None, aux_logits=True) model.fc = nn.Linear(model.fc.in_features, 1) model.AuxLogits.fc = nn.Linear(model.AuxLogits.fc.in_features, 1) state = torch.load(args.checkpoint, map_location="cpu", weights_only=False) model.load_state_dict(state) model.aux_logits = False model = model.to(device) model.eval() all_logits, all_labels = [], [] with torch.no_grad(): for images, labels in loader: images = images.to(device) logits = model(images) if hasattr(logits, "logits"): logits = logits.logits all_logits.append(logits.cpu()) all_labels.append(labels.unsqueeze(1)) all_logits = torch.cat(all_logits) all_labels = torch.cat(all_labels) acc, prec, rec, f1 = compute_clf_metrics(all_logits, all_labels) print(f"\n── Classification results ────────────────────") print(f" Provider : {args.provider}") print(f" Checkpoint : {args.checkpoint}") print(f" Test images: {len(test_ds):,}") print(f" Accuracy : {acc:.4f}") print(f" Precision : {prec:.4f}") print(f" Recall : {rec:.4f}") print(f" F1 : {f1:.4f}") # ── Entry point ─────────────────────────────────────────────────────────────── if __name__ == "__main__": parser = argparse.ArgumentParser() sub = parser.add_subparsers(dest="task", required=True) for name in ["seg", "clf"]: p = sub.add_parser(name) p.add_argument("--provider", required=True, choices=["google", "ign"]) p.add_argument("--checkpoint", required=True) p.add_argument("--batch_size", type=int, default=8) p.add_argument("--img_size", type=int, default=400) p.add_argument("--num_workers", type=int, default=0) args = parser.parse_args() if args.task == "seg": eval_seg(args) else: eval_clf(args)