#!/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") parser.add_argument("--export-existing-best", action="store_true") args = parser.parse_args() if args.output_dir.exists() and not args.export_existing_best: 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}") device = torch.device(args.device) model = resnet18(weights=weights) model.fc = nn.Linear(model.fc.in_features, 1) model.to(device) if args.export_existing_best: state_path = args.output_dir / "best-state.pt" if not state_path.is_file(): raise RuntimeError(f"missing existing best state: {state_path}") model.load_state_dict(torch.load(state_path, map_location=device)) model.eval() torch.jit.script(model).save(str(args.output_dir / "proposal-classifier.torchscript.pt")) print(json.dumps({"status": "exported_existing_best", "model": str(state_path)})) return 0 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) 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())