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