Add hard-negative proposal classifier pipeline
GeoIntel release gates / Compile, test, contracts and builds (push) Canceled after 0s
GeoIntel release gates / Python and npm vulnerability policy (push) Canceled after 0s
GeoIntel release gates / GIS image, SBOM and container scan (push) Canceled after 0s

This commit is contained in:
Jens
2026-07-30 01:38:39 +02:00
parent 3e66c78694
commit 68099b4a4e
3 changed files with 302 additions and 0 deletions
@@ -0,0 +1,96 @@
#!/usr/bin/env python3
"""Train and export the binary building-proposal filter."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
def binary_metrics(scores: list[float], labels: list[int], threshold: float = 0.5) -> dict[str, float | int]:
tp = sum(score >= threshold and label == 1 for score, label in zip(scores, labels, strict=True))
fp = sum(score >= threshold and label == 0 for score, label in zip(scores, labels, strict=True))
fn = sum(score < threshold and label == 1 for score, label in zip(scores, labels, strict=True))
precision = tp / (tp + fp) if tp + fp else 1.0
recall = tp / (tp + fn) if tp + fn else 1.0
return {"tp": tp, "fp": fp, "fn": fn, "precision": precision, "recall": recall,
"f1": 2 * precision * recall / (precision + recall) if precision + recall else 0.0}
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--dataset-dir", type=Path, required=True)
parser.add_argument("--output-dir", type=Path, required=True)
parser.add_argument("--epochs", type=int, default=12)
parser.add_argument("--batch", type=int, default=64)
parser.add_argument("--lr", type=float, default=1e-4)
parser.add_argument("--device", default="cuda:0")
args = parser.parse_args()
if args.output_dir.exists():
parser.error(f"output already exists: {args.output_dir}")
import torch
from torch import nn
from torch.utils.data import DataLoader
from torchvision.datasets import ImageFolder
from torchvision.models import ResNet18_Weights, resnet18
weights = ResNet18_Weights.DEFAULT
transform = weights.transforms()
train_ds = ImageFolder(args.dataset_dir / "train", transform=transform)
val_ds = ImageFolder(args.dataset_dir / "val", transform=transform)
if train_ds.class_to_idx != {"negative": 0, "positive": 1}:
raise RuntimeError(f"unexpected class order: {train_ds.class_to_idx}")
train_loader = DataLoader(train_ds, batch_size=args.batch, shuffle=True, num_workers=0)
val_loader = DataLoader(val_ds, batch_size=args.batch, shuffle=False, num_workers=0)
device = torch.device(args.device)
model = resnet18(weights=weights)
model.fc = nn.Linear(model.fc.in_features, 1)
model.to(device)
positives = sum(label == 1 for _path, label in train_ds.samples)
negatives = len(train_ds) - positives
loss_fn = nn.BCEWithLogitsLoss(pos_weight=torch.tensor([negatives / positives], device=device))
optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=1e-4)
args.output_dir.mkdir(parents=True)
history = []
best_f1 = -1.0
for epoch in range(1, args.epochs + 1):
model.train()
train_loss = 0.0
for images, labels in train_loader:
images, labels = images.to(device), labels.float().to(device)
optimizer.zero_grad(set_to_none=True)
logits = model(images).flatten()
loss = loss_fn(logits, labels)
loss.backward()
optimizer.step()
train_loss += float(loss) * len(images)
model.eval()
scores: list[float] = []
labels_out: list[int] = []
with torch.inference_mode():
for images, labels in val_loader:
scores.extend(torch.sigmoid(model(images.to(device)).flatten()).cpu().tolist())
labels_out.extend(labels.tolist())
metric = binary_metrics(scores, labels_out)
row = {"epoch": epoch, "train_loss": train_loss / len(train_ds), **metric}
history.append(row)
print(json.dumps(row), flush=True)
if float(metric["f1"]) > best_f1:
best_f1 = float(metric["f1"])
torch.save(model.state_dict(), args.output_dir / "best-state.pt")
model.load_state_dict(torch.load(args.output_dir / "best-state.pt", map_location=device))
model.eval()
scripted = torch.jit.script(model)
scripted.save(str(args.output_dir / "proposal-classifier.torchscript.pt"))
report = {"schema_version": 1, "status": "ok", "classes": train_ds.class_to_idx,
"train_count": len(train_ds), "validation_count": len(val_ds),
"best_validation_f1": best_f1, "history": history,
"model": str(args.output_dir / "proposal-classifier.torchscript.pt")}
(args.output_dir / "training-report.json").write_text(json.dumps(report, indent=2), encoding="utf-8")
return 0
if __name__ == "__main__":
raise SystemExit(main())