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())