File size: 7,565 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 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 | 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
|