CD-Models / datasets /cd_dataset.py
Dineth Perera
Publish tested dataset winners and benchmark rankings
ce209f5
Raw
History Blame Contribute Delete
7.57 kB
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