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: # pragma: no cover - torchvision import errors are environment-specific. 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