File size: 6,749 Bytes
a2aee5d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 | """
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)
|