CD-Models / train /fc_adapter.py
Dineth Perera
Publish tested dataset winners and benchmark rankings
ce209f5
Raw
History Blame Contribute Delete
13.8 kB
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())