CD-Models / utils /config_loader.py
Dineth Perera
Publish tested dataset winners and benchmark rankings
ce209f5
Raw
History Blame Contribute Delete
2.98 kB
from __future__ import annotations
import os
from copy import deepcopy
from pathlib import Path
from typing import Any
import yaml
from utils.dataset_cache import apply_io_defaults, resolve_effective_dataset_root
ROOT = Path(__file__).resolve().parents[1]
DEFAULT_DATA_ROOT = ROOT.parent / "Datasets"
FALLBACK_DATA_ROOTS = (ROOT.parent / "mamba-cd" / "datasets",)
def _read_yaml(path: Path) -> dict[str, Any]:
if not path.exists():
raise FileNotFoundError(f"Config file not found: {path}")
with path.open("r", encoding="utf-8") as f:
return yaml.safe_load(f) or {}
def _expand_data_root(value: str) -> str:
unresolved = value
if "${DATA_ROOT}" in value:
data_root = os.environ.get("DATA_ROOT")
if not data_root and DEFAULT_DATA_ROOT.is_dir():
data_root = str(DEFAULT_DATA_ROOT)
if not data_root:
raise EnvironmentError(
"DATA_ROOT is not set. Set DATA_ROOT to the parent directory containing "
"LEVIR-CD-plus-256, WHU-CD, WildFireS2, and other dataset folders."
)
value = value.replace("${DATA_ROOT}", data_root)
expanded = Path(value).expanduser().resolve()
if expanded.exists() or "${DATA_ROOT}" not in unresolved:
return str(expanded)
suffix = unresolved.split("${DATA_ROOT}", 1)[1].lstrip("/\\")
for root in FALLBACK_DATA_ROOTS:
candidate = (root / suffix).expanduser().resolve()
if candidate.exists():
return str(candidate)
return str(expanded)
def list_available_datasets() -> list[str]:
configs_dir = ROOT / "configs" / "datasets"
return sorted(path.stem for path in configs_dir.glob("*.yaml"))
def load_dataset_config(dataset_name: str) -> dict[str, Any]:
path = Path(dataset_name)
if not path.suffix:
path = ROOT / "configs" / "datasets" / f"{dataset_name}.yaml"
elif not path.is_absolute():
path = ROOT / path
cfg = _read_yaml(path)
if "data_root" not in cfg:
raise KeyError(f"{path} must define data_root")
cfg["data_root"] = _expand_data_root(str(cfg["data_root"]))
cfg["_config_path"] = str(path)
cfg["_dataset_name"] = path.stem
effective = resolve_effective_dataset_root(path.stem, cfg)
cfg.update(effective)
apply_io_defaults(cfg)
return cfg
def load_model_config(model_name: str) -> dict[str, Any]:
path = Path(model_name)
if not path.suffix:
path = ROOT / "configs" / "models" / f"{model_name}.yaml"
elif not path.is_absolute():
path = ROOT / path
return _read_yaml(path)
def merge_configs(model_cfg: dict[str, Any], dataset_cfg: dict[str, Any]) -> dict[str, Any]:
merged = deepcopy(model_cfg)
merged["dataset"] = deepcopy(dataset_cfg)
for key in ("batch_size", "num_workers", "img_size", "channels", "num_classes", "ignore_index"):
if key in dataset_cfg:
merged[key] = deepcopy(dataset_cfg[key])
return merged