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