| from __future__ import annotations |
|
|
| from pathlib import Path |
| from typing import Dict, Iterable, Optional |
|
|
| import numpy as np |
| import torch |
| from PIL import Image |
| from torch.utils.data import DataLoader, Dataset |
| from utils.dataset_cache import dataloader_kwargs, print_dataloader_policy |
|
|
| try: |
| from torchvision import transforms |
| except Exception: |
| transforms = None |
|
|
|
|
| IMG_EXTS = {".png", ".jpg", ".jpeg", ".tif", ".tiff", ".bmp"} |
| MASK_DIR_CANDIDATES = ("Mask", "mask", "label", "labels", "OUT", "gt") |
|
|
|
|
| def _scan(folder: Path) -> Dict[str, Path]: |
| return { |
| p.stem: p |
| for p in sorted(folder.iterdir()) |
| if p.is_file() and not p.name.startswith(".") and p.suffix.lower() in IMG_EXTS |
| } |
|
|
|
|
| def _resolve(root: Path, value: str | Path) -> Path: |
| p = Path(value) |
| return p if p.is_absolute() else root / p |
|
|
|
|
| def _configured_dirs(root: Path, split: str, cfg: dict) -> Optional[tuple[Path, Path, Path]]: |
| if "splits" in cfg: |
| split_dir = root / cfg.get("splits", {}).get(split, split) |
| a_dir = split_dir / cfg.get("image_a_folder", "A") |
| b_dir = split_dir / cfg.get("image_b_folder", "B") |
| m_dir = split_dir / cfg.get("mask_folder", "label") |
| return a_dir, b_dir, m_dir |
| source = cfg.get("split", {}).get("source_folders", {}) |
| a_rel = source.get(f"{split}_a") or source.get("a") |
| b_rel = source.get(f"{split}_b") or source.get("b") |
| m_rel = source.get(f"{split}_mask") or source.get("mask") |
| if not (a_rel and b_rel and m_rel): |
| return None |
| return _resolve(root, a_rel), _resolve(root, b_rel), _resolve(root, m_rel) |
|
|
|
|
| def _mask_dir(split_dir: Path) -> Path: |
| for name in MASK_DIR_CANDIDATES: |
| candidate = split_dir / name |
| if candidate.is_dir(): |
| return candidate |
| tried = ", ".join(MASK_DIR_CANDIDATES) |
| raise FileNotFoundError(f"No mask directory under {split_dir}; tried {tried}") |
|
|
|
|
| class CDDataset(Dataset): |
| """Configurable binary change-detection dataset. |
| |
| This is the shared root loader for CD-Models. It intentionally supports the |
| folder conventions found in mamba-cd and the nested research repos instead |
| of adding another WildFire-specific loader. |
| """ |
|
|
| def __init__( |
| self, |
| root: str | Path, |
| split: str, |
| cfg: Optional[dict] = None, |
| image_size: Optional[int] = None, |
| normalize: bool = True, |
| return_format: str = "dict", |
| ) -> None: |
| self.root = Path(root) |
| self.split = split |
| self.cfg = cfg or {} |
| self.return_format = return_format |
| ds_cfg = self.cfg.get("dataset", self.cfg) |
| self.image_size = int(image_size or ds_cfg.get("image_size", ds_cfg.get("img_size", 256))) |
| self.threshold = int(ds_cfg.get("binary_threshold", ds_cfg.get("mask_threshold", 127))) |
| self.mean = ds_cfg.get("mean", ds_cfg.get("mean_a", [0.485, 0.456, 0.406])) |
| self.std = ds_cfg.get("std", ds_cfg.get("std_a", [0.229, 0.224, 0.225])) |
| if not self.root.is_dir(): |
| dataset_name = ds_cfg.get("_dataset_name", ds_cfg.get("name", "unknown")) |
| raise FileNotFoundError( |
| f"\n{'=' * 60}\n" |
| f"Dataset root directory not found:\n {self.root}\n\n" |
| f"For dataset '{dataset_name}', either:\n" |
| f" 1. Download the dataset and place it at the above path\n" |
| f" 2. Update DATA_ROOT: export DATA_ROOT=/correct/path\n" |
| f" 3. Edit configs/datasets/{dataset_name}.yaml\n" |
| f"{'=' * 60}" |
| ) |
|
|
| configured = _configured_dirs(self.root, split, self.cfg) |
| if configured is None: |
| split_dir = self.root / split |
| a_dir = split_dir / "A" |
| b_dir = split_dir / "B" |
| m_dir = _mask_dir(split_dir) |
| else: |
| a_dir, b_dir, m_dir = configured |
|
|
| for label, folder in (("A", a_dir), ("B", b_dir), ("mask", m_dir)): |
| if not folder.is_dir(): |
| raise FileNotFoundError(f"{label} directory not found: {folder}") |
|
|
| a_files = _scan(a_dir) |
| b_files = _scan(b_dir) |
| m_files = _scan(m_dir) |
| stems = sorted(set(a_files) & set(b_files) & set(m_files)) |
| if not stems: |
| raise RuntimeError(f"No matched A/B/mask triplets found for {self.root} split {split}") |
| self.samples = [(a_files[s], b_files[s], m_files[s], s) for s in stems] |
|
|
| if transforms is None: |
| self.image_tf = None |
| else: |
| ops: list = [transforms.Resize((self.image_size, self.image_size)), transforms.ToTensor()] |
| if normalize: |
| ops.append(transforms.Normalize(mean=self.mean, std=self.std)) |
| self.image_tf = transforms.Compose(ops) |
|
|
| def __len__(self) -> int: |
| return len(self.samples) |
|
|
| def _image(self, path: Path) -> torch.Tensor: |
| img = Image.open(path).convert("RGB") |
| if self.image_tf is not None: |
| return self.image_tf(img) |
| img = img.resize((self.image_size, self.image_size), Image.BILINEAR) |
| arr = np.array(img, dtype=np.float32) / 255.0 |
| return torch.from_numpy(arr).permute(2, 0, 1) |
|
|
| def _mask(self, path: Path) -> torch.Tensor: |
| mask = Image.open(path).convert("L").resize((self.image_size, self.image_size), Image.NEAREST) |
| arr = np.array(mask, dtype=np.uint8) |
| if arr.max() <= 1: |
| bin_mask = (arr > 0).astype(np.float32) |
| else: |
| bin_mask = (arr > self.threshold).astype(np.float32) |
| return torch.from_numpy(bin_mask).unsqueeze(0) |
|
|
| def __getitem__(self, index: int): |
| a_path, b_path, mask_path, name = self.samples[index] |
| a = self._image(a_path) |
| b = self._image(b_path) |
| mask = self._mask(mask_path) |
| if self.return_format == "tuple": |
| return a, b, mask, name |
| if self.return_format == "legacy": |
| return {"A": a, "B": b, "L": mask.squeeze(0).long(), "name": name} |
| return {"a": a, "b": b, "mask": mask, "name": name} |
|
|
|
|
| def build_dataloader( |
| cfg: dict, |
| split: str, |
| shuffle: Optional[bool] = None, |
| return_format: str = "dict", |
| ) -> DataLoader: |
| ds_cfg = cfg.get("dataset", cfg) |
| dataset = CDDataset( |
| ds_cfg.get("root", ds_cfg["data_root"]), |
| split=split, |
| cfg=cfg, |
| image_size=int(ds_cfg.get("image_size", ds_cfg.get("img_size", 256))), |
| return_format=return_format, |
| ) |
| batch_size = int(ds_cfg.get("batch_size", 8)) |
| num_workers = int(ds_cfg.get("num_workers", 4)) |
| if shuffle is None: |
| shuffle = split == "train" |
| print_dataloader_policy({**ds_cfg, "num_workers": num_workers}, torch.cuda.is_available()) |
| return DataLoader( |
| dataset, |
| batch_size=batch_size, |
| shuffle=shuffle, |
| **dataloader_kwargs({**ds_cfg, "num_workers": num_workers}, torch.cuda.is_available()), |
| drop_last=split == "train", |
| ) |
|
|
|
|
| def available_splits(cfg: dict) -> Iterable[str]: |
| ds_cfg = cfg.get("dataset", cfg) |
| root = Path(ds_cfg.get("root", ds_cfg["data_root"])) |
| if "splits" in cfg: |
| yield from cfg["splits"] |
| return |
| source = cfg.get("split", {}).get("source_folders", {}) |
| for split in ("train", "val", "test"): |
| if source.get(f"{split}_a") and source.get(f"{split}_b") and source.get(f"{split}_mask"): |
| yield split |
| elif (root / split).is_dir(): |
| yield split |
|
|