CD-Models / run_training.py
Dineth Perera
Publish tested dataset winners and benchmark rankings
ce209f5
Raw
History Blame Contribute Delete
9.38 kB
"""
run_training.py - Master training launcher for cd-models benchmark suite.
Usage examples:
python run_training.py --model changemamba --dataset levir_cd
python run_training.py --model bifa --dataset all
python run_training.py --model all --dataset wildfire_s2
python run_training.py --model all --dataset all
python run_training.py --model elgcnet --dataset whu_cd --resume
python run_training.py --model hanet --dataset levir_cd --eval-only
"""
from __future__ import annotations
import argparse
import json
import subprocess
import sys
from datetime import datetime, timezone
from pathlib import Path
import yaml
from utils.config_loader import load_dataset_config
from utils.dataset_cache import dataset_runtime_summary, print_dataloader_policy
from utils.gpu_utils import subprocess_gpu_env
from utils.results_writer import append_to_comparison_table
ROOT = Path(__file__).resolve().parent
SKIP_EXIT_CODE = 75
FAILED_SUBPROCESS_EXIT_CODE = 70
FAILED_NO_CHECKPOINT_EXIT_CODE = 71
FAILED_EVAL_EXIT_CODE = 72
WILDFIRE_ONLY_MODELS = set()
MAMBA_MODELS = {"cdmamba", "changemamba", "rsm_cd"}
GLOBAL_DATASET_EXCLUSIONS = {
"kate_cd": "KATE-CD training is disabled for the current benchmark sweep.",
"levir_cd": "Use levir_cd_test_as_val instead of the standard LEVIR-CD split.",
}
DATASET_EXCLUSIONS = {
"tinycd": {"wildfire_s2"},
}
def _load_datasets() -> list[str]:
return sorted(path.stem for path in (ROOT / "configs" / "datasets").glob("*.yaml"))
def _load_registry() -> dict:
with (ROOT / "configs" / "models" / "registry.yaml").open("r", encoding="utf-8") as f:
return yaml.safe_load(f)["models"]
def _parse_args() -> argparse.Namespace:
registry = _load_registry()
datasets = _load_datasets()
parser = argparse.ArgumentParser(description="Master training launcher for cd-models benchmark suite.")
parser.add_argument("--model", required=True, choices=sorted(registry) + ["all"])
parser.add_argument("--dataset", required=True, choices=datasets + ["all"])
parser.add_argument("--resume", action="store_true")
parser.add_argument("--eval-only", action="store_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("--dry-run", action="store_true")
parser.add_argument("--force", action="store_true")
parser.add_argument("--max-iters", type=int, default=None, help="Pass an iteration override to iteration-based wrappers.")
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 _log_run(row: dict) -> None:
path = ROOT / "results" / "training_log.jsonl"
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("a", encoding="utf-8") as f:
f.write(json.dumps(row, sort_keys=True) + "\n")
def _status_symbol(status: str) -> str:
return status
def _status_from_exit_code(code: int, dry_run: bool) -> str:
if code == SKIP_EXIT_CODE:
return "skipped_dataset_policy"
if dry_run:
return "dry-run" if code == 0 else "failed_subprocess"
if code == 0:
return "trained"
if code == FAILED_SUBPROCESS_EXIT_CODE:
return "failed_subprocess"
if code == FAILED_NO_CHECKPOINT_EXIT_CODE:
return "failed_no_checkpoint"
if code == FAILED_EVAL_EXIT_CODE:
return "failed_eval"
return "failed_subprocess"
def _dataset_root_exists(dataset: str) -> tuple[bool, str]:
cfg = load_dataset_config(dataset)
root = Path(cfg["data_root"])
print(f"[DATASET] {dataset_runtime_summary(cfg)}", flush=True)
print_dataloader_policy(cfg)
if cfg.get("io_warning"):
print(f"[DATASET-WARNING] {cfg['io_warning']}", flush=True)
return root.exists(), str(root)
def main() -> int:
args = _parse_args()
registry = _load_registry()
available_datasets = _load_datasets()
models = sorted(registry) if args.model == "all" else [args.model]
datasets = available_datasets if args.dataset == "all" else [args.dataset]
rows = []
dataset_roots = {}
for dataset in datasets:
try:
dataset_roots[dataset] = _dataset_root_exists(dataset)
except Exception as exc:
dataset_roots[dataset] = (False, f"config error: {exc}")
for model in models:
for dataset in datasets:
root_ok, root_note = dataset_roots[dataset]
if not root_ok:
print(f"[SKIP] {model}/{dataset}: dataset root is not available ({root_note}).", flush=True)
rows.append((model, dataset, "skipped_missing_dataset"))
continue
if dataset in GLOBAL_DATASET_EXCLUSIONS and model not in MAMBA_MODELS:
print(f"[SKIP] {model}/{dataset}: {GLOBAL_DATASET_EXCLUSIONS[dataset]}", flush=True)
rows.append((model, dataset, "skipped_dataset_policy"))
continue
if model in WILDFIRE_ONLY_MODELS and dataset != "wildfire_s2":
print(f"[SKIP] {model}/{dataset}: wrapper is wired only to WildFireS2 training.", flush=True)
rows.append((model, dataset, "skipped_dataset_policy"))
continue
excluded = DATASET_EXCLUSIONS.get(model, set())
if dataset in excluded:
print(f"[SKIP] {model}/{dataset} unsupported.", flush=True)
rows.append((model, dataset, "skipped_dataset_policy"))
continue
metrics_path = ROOT / "results" / model / dataset / "metrics_test.json"
if metrics_path.exists() and not args.force:
print(f"[SKIP] {model}/{dataset} already complete", flush=True)
rows.append((model, dataset, "skipped_complete"))
continue
script = ROOT / registry[model]["script"]
cmd = [sys.executable, str(script), "--dataset", dataset, "--gpu", str(args.gpu)]
if model.startswith("fc_"):
cmd.extend(["--model", model])
if args.resume:
cmd.append("--resume")
if args.eval_only:
cmd.append("--eval-only")
if args.epochs is not None:
cmd.extend(["--epochs", str(args.epochs)])
if args.batch_size is not None:
cmd.extend(["--batch-size", str(args.batch_size)])
if args.lr is not None:
cmd.extend(["--lr", str(args.lr)])
if args.dry_run:
cmd.append("--dry-run")
if args.force:
cmd.append("--force")
if args.max_iters is not None:
cmd.extend(["--max-iters", str(args.max_iters)])
if args.num_workers is not None:
cmd.extend(["--num-workers", str(args.num_workers)])
if args.prefetch_factor is not None:
cmd.extend(["--prefetch-factor", str(args.prefetch_factor)])
if args.persistent_workers is True:
cmd.append("--persistent-workers")
elif args.persistent_workers is False:
cmd.append("--no-persistent-workers")
if args.pin_memory is True:
cmd.append("--pin-memory")
elif args.pin_memory is False:
cmd.append("--no-pin-memory")
start = datetime.now(timezone.utc)
child_env, child_gpu = subprocess_gpu_env(args.gpu)
print(f"[GPU] requested physical GPU: {child_gpu.requested_gpu}", flush=True)
print(f"[GPU] launching with CUDA_VISIBLE_DEVICES={child_gpu.cuda_visible_devices}", flush=True)
print("[RUN]", model, dataset, " ".join(cmd), flush=True)
code = subprocess.run(cmd, cwd=ROOT, env=child_env, check=False).returncode
end = datetime.now(timezone.utc)
status = _status_from_exit_code(code, args.dry_run)
rows.append((model, dataset, status))
_log_run({
"model": model,
"dataset": dataset,
"command": cmd,
"start_time": start.isoformat(),
"end_time": end.isoformat(),
"exit_code": code,
"status": status,
})
append_to_comparison_table()
print("\nModel | Dataset | Status", flush=True)
print("--- | --- | ---", flush=True)
for model, dataset, status in rows:
print(f"{model} | {dataset} | {_status_symbol(status)}", flush=True)
ok_statuses = {"trained", "skipped_dataset_policy", "skipped_missing_dataset", "skipped_complete", "dry-run"}
return 0 if all(status in ok_statuses for _, _, status in rows) else 1
if __name__ == "__main__":
raise SystemExit(main())