| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import sys |
| from datetime import datetime, timezone |
| from pathlib import Path |
|
|
| import torch |
| import torch.nn as nn |
| from torch.utils.data import DataLoader |
|
|
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| REPO = ROOT / "model_repos" / "fully_convolutional_change_detection" |
| if str(ROOT) not in sys.path: |
| sys.path.insert(0, str(ROOT)) |
| if str(REPO) not in sys.path: |
| sys.path.insert(0, str(REPO)) |
|
|
| from datasets.cd_dataset import CDDataset |
| from siamunet_conc import SiamUnet_conc |
| from siamunet_diff import SiamUnet_diff |
| from unet import Unet |
| from utils.config_loader import load_dataset_config, load_model_config |
| from utils.dataset_cache import apply_dataloader_cli_overrides, dataloader_kwargs, dataloader_policy_lines, dataset_runtime_summary, print_dataloader_policy |
| from utils.gpu_utils import print_gpu_diagnostics, resolve_gpu |
| from utils.metrics import BinaryMetrics |
| from utils.model_adapters import get_model_adapter |
| from utils.results_writer import append_to_comparison_table, save_metrics |
| from utils.unified_evaluator import evaluate_with_adapter |
|
|
|
|
| def _parse_args() -> argparse.Namespace: |
| parser = argparse.ArgumentParser(description="Train FC-EF / FC-Siam variants on a CD-Models dataset.") |
| parser.add_argument("--model", required=True, choices=["fc_ef", "fc_siam_conc", "fc_siam_diff"]) |
| parser.add_argument("--dataset", required=True) |
| parser.add_argument("--epochs", type=int, default=None) |
| parser.add_argument("--batch-size", type=int, default=None) |
| parser.add_argument("--lr", type=float, default=None) |
| parser.add_argument("--gpu", default="0") |
| parser.add_argument("--resume", action="store_true") |
| parser.add_argument("--force", action="store_true") |
| parser.add_argument("--eval-only", action="store_true") |
| parser.add_argument("--smoke-test", action="store_true") |
| parser.add_argument("--dry-run", action="store_true") |
| parser.add_argument("--output-dir", default=None) |
| parser.add_argument("--allow-missing-profilers", action="store_true") |
| parser.add_argument("--num-workers", type=int, default=None) |
| parser.add_argument("--prefetch-factor", type=int, default=None) |
| parser.set_defaults(persistent_workers=None, pin_memory=None) |
| parser.add_argument("--persistent-workers", dest="persistent_workers", action="store_true") |
| parser.add_argument("--no-persistent-workers", dest="persistent_workers", action="store_false") |
| parser.add_argument("--pin-memory", dest="pin_memory", action="store_true") |
| parser.add_argument("--no-pin-memory", dest="pin_memory", action="store_false") |
| return parser.parse_args() |
|
|
|
|
| def _build_model(model_name: str) -> nn.Module: |
| if model_name == "fc_ef": |
| return Unet(input_nbr=6, label_nbr=2) |
| if model_name == "fc_siam_conc": |
| return SiamUnet_conc(input_nbr=3, label_nbr=2) |
| return SiamUnet_diff(input_nbr=3, label_nbr=2) |
|
|
|
|
| def _metrics(log_probs: torch.Tensor, target: torch.Tensor) -> tuple[int, int, int, int]: |
| pred = log_probs.argmax(dim=1) |
| tp = ((pred == 1) & (target == 1)).sum().item() |
| fp = ((pred == 1) & (target == 0)).sum().item() |
| fn = ((pred == 0) & (target == 1)).sum().item() |
| tn = ((pred == 0) & (target == 0)).sum().item() |
| return tp, fp, fn, tn |
|
|
|
|
| def _score(tp: int, fp: int, fn: int, tn: int) -> dict[str, float]: |
| eps = 1e-8 |
| precision = tp / (tp + fp + eps) |
| recall = tp / (tp + fn + eps) |
| f1 = 2 * precision * recall / (precision + recall + eps) |
| iou = tp / (tp + fp + fn + eps) |
| acc = (tp + tn) / (tp + fp + fn + tn + eps) |
| return {"f1": f1, "iou": iou, "precision": precision, "recall": recall, "accuracy": acc} |
|
|
|
|
| def _forward(model: nn.Module, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: |
| return model(a, b) |
|
|
|
|
| def _fc_loss(model_name: str, log_probs: torch.Tensor, target: torch.Tensor) -> torch.Tensor: |
| weight = None |
| if model_name == "fc_siam_diff": |
| pos = target.eq(1).sum().float() |
| neg = target.eq(0).sum().float() |
| if bool(pos.item() > 0): |
| pos_weight = (neg / pos.clamp_min(1.0)).clamp(min=1.0, max=50.0) |
| weight = torch.stack([torch.ones_like(pos_weight), pos_weight]).to(log_probs.device) |
| return nn.functional.nll_loss(log_probs.float(), target, weight=weight) |
|
|
|
|
| def _load_checkpoint_if_requested(model: nn.Module, optimizer: torch.optim.Optimizer, path: Path) -> tuple[int, float]: |
| if not path.exists(): |
| raise FileNotFoundError(f"Cannot resume; checkpoint not found: {path}") |
| checkpoint = torch.load(path, map_location="cpu") |
| if not isinstance(checkpoint, dict) or "model_state_dict" not in checkpoint: |
| raise RuntimeError(f"Checkpoint {path} is not a CD-Models FC checkpoint.") |
| model.load_state_dict(checkpoint["model_state_dict"], strict=True) |
| if "optimizer_state_dict" in checkpoint: |
| optimizer.load_state_dict(checkpoint["optimizer_state_dict"]) |
| return int(checkpoint.get("epoch", 0)), float(checkpoint.get("best_val_f1", -1.0)) |
|
|
|
|
| def _run_smoke_test(args: argparse.Namespace) -> int: |
| gpu_resolution = resolve_gpu(args.gpu) |
| print_gpu_diagnostics(gpu_resolution) |
| dataset_cfg = load_dataset_config(args.dataset) |
| apply_dataloader_cli_overrides(dataset_cfg, args) |
| print_dataloader_policy(dataset_cfg, torch.cuda.is_available()) |
| ds = CDDataset(dataset_cfg["data_root"], "test", cfg=dataset_cfg, return_format="tuple") |
| a, b, mask, name = ds[0] |
| model = _build_model(args.model) |
| with torch.inference_mode(): |
| output = model(a.unsqueeze(0), b.unsqueeze(0)) |
| metrics = BinaryMetrics(threshold=float(dataset_cfg.get("eval", {}).get("threshold", 0.5))) |
| metrics.update(output, mask.unsqueeze(0)) |
| param_count = sum(p.numel() for p in model.parameters()) |
| smoke_dir = ROOT / "results" / args.model / dataset_cfg["name"] / "smoke_test" |
| smoke_dir.mkdir(parents=True, exist_ok=True) |
| from utils.metrics import normalize_binary_prediction |
| from utils.qualitative import save_binary_prediction |
| from utils.profiling import ProfilingUnavailable, count_flops |
|
|
| pred, _ = normalize_binary_prediction(output) |
| save_binary_prediction(pred[0], smoke_dir / f"{name}_pred.png") |
| flops_error = None |
| try: |
| count_flops( |
| model, |
| lambda: ( |
| torch.zeros(1, 3, int(dataset_cfg.get("img_size", 256)), int(dataset_cfg.get("img_size", 256))), |
| torch.zeros(1, 3, int(dataset_cfg.get("img_size", 256)), int(dataset_cfg.get("img_size", 256))), |
| ), |
| torch.device("cpu"), |
| ) |
| except ProfilingUnavailable as exc: |
| flops_error = str(exc) |
| print( |
| f"[SMOKE] {args.model}/{dataset_cfg['name']} sample={name} output={tuple(output.shape)} " |
| f"params={param_count} metrics={metrics.compute()} flops_error={flops_error}" |
| ) |
| return 0 |
|
|
|
|
| def main() -> int: |
| args = _parse_args() |
| if args.smoke_test: |
| return _run_smoke_test(args) |
| gpu_resolution = resolve_gpu(args.gpu) |
| print_gpu_diagnostics(gpu_resolution) |
| device = torch.device(gpu_resolution.local_device) |
|
|
| dataset_cfg = load_dataset_config(args.dataset) |
| apply_dataloader_cli_overrides(dataset_cfg, args) |
| model_cfg = load_model_config(args.model) |
| if args.dry_run: |
| print(f"[DRY-RUN] fc model={args.model} dataset={dataset_cfg['name']} root={dataset_cfg['data_root']}") |
| print(f"[DATASET] {dataset_runtime_summary(dataset_cfg)}") |
| print_dataloader_policy(dataset_cfg, torch.cuda.is_available()) |
| return 0 |
| train_ds = CDDataset(dataset_cfg["data_root"], "train", cfg=dataset_cfg, return_format="tuple") |
| val_ds = CDDataset(dataset_cfg["data_root"], "val", cfg=dataset_cfg, return_format="tuple") |
| batch_size = int(args.batch_size or dataset_cfg.get("batch_size", 8)) |
| loader_kwargs = dataloader_kwargs(dataset_cfg, torch.cuda.is_available()) |
| print(f"[DATASET] {dataset_runtime_summary(dataset_cfg)}") |
| print_dataloader_policy(dataset_cfg, torch.cuda.is_available()) |
| if dataset_cfg.get("io_warning"): |
| print(f"[DATASET-WARNING] {dataset_cfg['io_warning']}") |
| train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, **loader_kwargs, drop_last=True) |
| val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, **loader_kwargs) |
|
|
| model = _build_model(args.model).to(device) |
| optimizer = torch.optim.Adam( |
| model.parameters(), |
| lr=float(args.lr or model_cfg.get("lr", 1e-4)), |
| weight_decay=float(model_cfg.get("weight_decay", 0.0) or 0.0), |
| ) |
|
|
| out_dir = Path(args.output_dir) if args.output_dir else ROOT / "results" / args.model / dataset_cfg["name"] |
| if not out_dir.is_absolute(): |
| out_dir = ROOT / out_dir |
| ckpt_dir = out_dir / "checkpoints" |
| log_dir = out_dir / "logs" |
| ckpt_dir.mkdir(parents=True, exist_ok=True) |
| log_dir.mkdir(parents=True, exist_ok=True) |
| log_path = log_dir / "train.log" |
| best_f1 = -1.0 |
| best_epoch = 0 |
| start_epoch = 1 |
| latest_path = ckpt_dir / "latest.pth" |
| best_path = ckpt_dir / "best_model.pth" |
|
|
| if args.resume: |
| last_epoch, best_f1 = _load_checkpoint_if_requested(model, optimizer, latest_path) |
| start_epoch = last_epoch + 1 |
| best_epoch = int(last_epoch) |
|
|
| if args.eval_only: |
| _, code = evaluate_with_adapter( |
| model_name=args.model, |
| dataset_cfg=dataset_cfg, |
| model_config=model_cfg, |
| adapter=get_model_adapter(args.model), |
| checkpoint_path=best_path, |
| device=device, |
| strict_profiling=not args.allow_missing_profilers, |
| output_dir=out_dir, |
| ) |
| return code |
|
|
| with log_path.open("w", encoding="utf-8") as log: |
| log.write(f"# {args.model} {dataset_cfg['name']} training\n") |
| log.write(f"# Dataset: {dataset_runtime_summary(dataset_cfg)}\n") |
| for line in dataloader_policy_lines(dataset_cfg, torch.cuda.is_available()): |
| log.write(line + "\n") |
| epochs = int(args.epochs or model_cfg.get("num_epochs", 200)) |
| train_start = datetime.now(timezone.utc) |
| for epoch in range(start_epoch, epochs + 1): |
| model.train() |
| total_loss = 0.0 |
| n_batches = 0 |
| for a, b, mask, _ in train_loader: |
| a = a.to(device) |
| b = b.to(device) |
| target = mask.squeeze(1).long().to(device) |
| optimizer.zero_grad() |
| loss = _fc_loss(args.model, model(a, b), target) |
| loss.backward() |
| optimizer.step() |
| total_loss += loss.item() |
| n_batches += 1 |
|
|
| model.eval() |
| val_metrics = BinaryMetrics(threshold=float(dataset_cfg.get("eval", {}).get("threshold", 0.5))) |
| val_start = datetime.now(timezone.utc) |
| with torch.no_grad(): |
| for a, b, mask, _ in val_loader: |
| scores = model(a.to(device), b.to(device)) |
| val_metrics.update(scores.detach().cpu(), mask) |
| val_end = datetime.now(timezone.utc) |
| scores = val_metrics.compute() |
| train_loss = total_loss / max(n_batches, 1) |
| is_best = scores["f1"] > best_f1 |
| checkpoint = { |
| "model": args.model, |
| "dataset": dataset_cfg["name"], |
| "epoch": epoch, |
| "best_epoch": best_epoch, |
| "best_val_f1": best_f1, |
| "model_state_dict": model.state_dict(), |
| "optimizer_state_dict": optimizer.state_dict(), |
| } |
| if is_best: |
| best_f1 = scores["f1"] |
| best_epoch = epoch |
| checkpoint["best_epoch"] = best_epoch |
| checkpoint["best_val_f1"] = best_f1 |
| torch.save(checkpoint, best_path) |
| torch.save(checkpoint, latest_path) |
| val_payload = dict(scores) |
| val_payload.update({ |
| "model": args.model, |
| "dataset": dataset_cfg["name"], |
| "split": "val", |
| "epoch": epoch, |
| "best_epoch": best_epoch, |
| "checkpoint": str(best_path if is_best else latest_path), |
| "validation_time": (val_end - val_start).total_seconds(), |
| "timestamp": val_end.isoformat(), |
| "status": "complete", |
| }) |
| save_metrics(args.model, dataset_cfg["name"], "val", val_payload) |
| line = ( |
| f"[Epoch {epoch:03d}/{int(model_cfg.get('num_epochs', 200)):03d}] " |
| f"train_loss={train_loss:.4f} | val_F1={scores['f1']:.4f} " |
| f"val_IoU={scores['iou']:.4f} val_Prec={scores['precision']:.4f} " |
| f"val_Rec={scores['recall']:.4f} val_Acc={scores['accuracy']:.4f}" |
| f"{' | BEST' if is_best else ''}" |
| ) |
| print(line, flush=True) |
| log.write(line + "\n") |
| log.flush() |
|
|
| train_end = datetime.now(timezone.utc) |
| test_metrics, code = evaluate_with_adapter( |
| model_name=args.model, |
| dataset_cfg=dataset_cfg, |
| model_config=model_cfg, |
| adapter=get_model_adapter(args.model), |
| checkpoint_path=best_path, |
| device=device, |
| strict_profiling=not args.allow_missing_profilers, |
| output_dir=out_dir, |
| ) |
| test_metrics["best_epoch"] = best_epoch |
| test_metrics["training_time"] = (train_end - train_start).total_seconds() |
| test_metrics["best_val_f1"] = best_f1 |
| save_metrics(args.model, dataset_cfg["name"], "test", test_metrics) |
| append_to_comparison_table() |
| return code |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main()) |
|
|