File size: 2,983 Bytes
ce209f5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
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