bdappv-models / train /evaluate.py
gabrielkasmi's picture
Upload 9 files
a2aee5d verified
Raw History Blame Contribute Delete
6.75 kB
"""
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)