dataeval-flow 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- dataeval_flow/__init__.py +93 -0
- dataeval_flow/__main__.py +149 -0
- dataeval_flow/_app/__init__.py +5 -0
- dataeval_flow/_app/_model/__init__.py +5 -0
- dataeval_flow/_app/_model/_coerce.py +126 -0
- dataeval_flow/_app/_model/_discover.py +171 -0
- dataeval_flow/_app/_model/_execution.py +108 -0
- dataeval_flow/_app/_model/_introspect.py +280 -0
- dataeval_flow/_app/_model/_item.py +213 -0
- dataeval_flow/_app/_model/_registry.py +255 -0
- dataeval_flow/_app/_model/_state.py +322 -0
- dataeval_flow/_app/_model/_undo.py +61 -0
- dataeval_flow/_app/_panes/__init__.py +35 -0
- dataeval_flow/_app/_panes/_config_pane.py +173 -0
- dataeval_flow/_app/_panes/_result_pane.py +125 -0
- dataeval_flow/_app/_panes/_task_pane.py +91 -0
- dataeval_flow/_app/_panes/_widgets.py +111 -0
- dataeval_flow/_app/_screens/__init__.py +25 -0
- dataeval_flow/_app/_screens/_base.py +242 -0
- dataeval_flow/_app/_screens/_detail.py +333 -0
- dataeval_flow/_app/_screens/_model.py +102 -0
- dataeval_flow/_app/_screens/_params.py +80 -0
- dataeval_flow/_app/_screens/_pathpicker.py +68 -0
- dataeval_flow/_app/_screens/_section.py +621 -0
- dataeval_flow/_app/_screens/_settings.py +183 -0
- dataeval_flow/_app/_viewmodel/__init__.py +15 -0
- dataeval_flow/_app/_viewmodel/_builder_vm.py +272 -0
- dataeval_flow/_app/_viewmodel/_model_vm.py +70 -0
- dataeval_flow/_app/_viewmodel/_rendering.py +189 -0
- dataeval_flow/_app/_viewmodel/_result_vm.py +210 -0
- dataeval_flow/_app/_viewmodel/_section_vm.py +224 -0
- dataeval_flow/_app/app.py +742 -0
- dataeval_flow/_app/cli.py +592 -0
- dataeval_flow/_logging.py +102 -0
- dataeval_flow/cache.py +1355 -0
- dataeval_flow/config/__init__.py +80 -0
- dataeval_flow/config/_loader.py +79 -0
- dataeval_flow/config/_merge.py +92 -0
- dataeval_flow/config/_models.py +115 -0
- dataeval_flow/config/_paths.py +85 -0
- dataeval_flow/config/schemas/__init__.py +112 -0
- dataeval_flow/config/schemas/_dataset.py +111 -0
- dataeval_flow/config/schemas/_extractor.py +119 -0
- dataeval_flow/config/schemas/_metadata.py +28 -0
- dataeval_flow/config/schemas/_preprocessor.py +18 -0
- dataeval_flow/config/schemas/_selection.py +100 -0
- dataeval_flow/config/schemas/_task.py +89 -0
- dataeval_flow/config/schemas/_workflow.py +135 -0
- dataeval_flow/dataset.py +635 -0
- dataeval_flow/embeddings.py +135 -0
- dataeval_flow/metadata.py +48 -0
- dataeval_flow/preprocessing.py +141 -0
- dataeval_flow/py.typed +0 -0
- dataeval_flow/runner.py +118 -0
- dataeval_flow/selection.py +50 -0
- dataeval_flow/workflow/__init__.py +328 -0
- dataeval_flow/workflow/_text_report.py +511 -0
- dataeval_flow/workflow/base.py +69 -0
- dataeval_flow/workflow/orchestrator.py +454 -0
- dataeval_flow/workflows/__init__.py +1 -0
- dataeval_flow/workflows/analysis/__init__.py +38 -0
- dataeval_flow/workflows/analysis/outputs.py +202 -0
- dataeval_flow/workflows/analysis/params.py +114 -0
- dataeval_flow/workflows/analysis/workflow.py +1313 -0
- dataeval_flow/workflows/cleaning/__init__.py +23 -0
- dataeval_flow/workflows/cleaning/outputs.py +200 -0
- dataeval_flow/workflows/cleaning/params.py +160 -0
- dataeval_flow/workflows/cleaning/report.py +304 -0
- dataeval_flow/workflows/cleaning/workflow.py +794 -0
- dataeval_flow/workflows/drift/__init__.py +1 -0
- dataeval_flow/workflows/drift/outputs.py +144 -0
- dataeval_flow/workflows/drift/params.py +332 -0
- dataeval_flow/workflows/drift/report.py +201 -0
- dataeval_flow/workflows/drift/workflow.py +647 -0
- dataeval_flow/workflows/ood/__init__.py +1 -0
- dataeval_flow/workflows/ood/outputs.py +134 -0
- dataeval_flow/workflows/ood/params.py +161 -0
- dataeval_flow/workflows/ood/report.py +311 -0
- dataeval_flow/workflows/ood/workflow.py +728 -0
- dataeval_flow/workflows/prioritization/__init__.py +1 -0
- dataeval_flow/workflows/prioritization/outputs.py +122 -0
- dataeval_flow/workflows/prioritization/params.py +124 -0
- dataeval_flow/workflows/prioritization/report.py +117 -0
- dataeval_flow/workflows/prioritization/workflow.py +587 -0
- dataeval_flow/workflows/splitting/__init__.py +25 -0
- dataeval_flow/workflows/splitting/outputs.py +101 -0
- dataeval_flow/workflows/splitting/params.py +61 -0
- dataeval_flow/workflows/splitting/report.py +485 -0
- dataeval_flow/workflows/splitting/workflow.py +371 -0
- dataeval_flow-0.1.0.dist-info/METADATA +305 -0
- dataeval_flow-0.1.0.dist-info/RECORD +94 -0
- dataeval_flow-0.1.0.dist-info/WHEEL +4 -0
- dataeval_flow-0.1.0.dist-info/entry_points.txt +2 -0
- dataeval_flow-0.1.0.dist-info/licenses/LICENSE +21 -0
dataeval_flow/dataset.py
ADDED
|
@@ -0,0 +1,635 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
"""Dataset loading utilities - standalone library functions.
|
|
3
|
+
|
|
4
|
+
All functions accept parameters. No hardcoded container paths.
|
|
5
|
+
|
|
6
|
+
This module provides functions for loading datasets in HuggingFace, image folder,
|
|
7
|
+
COCO, YOLO, and torchvision formats, converting them to MAITE-compatible objects for evaluation.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
import logging
|
|
11
|
+
from dataclasses import dataclass
|
|
12
|
+
from functools import lru_cache
|
|
13
|
+
from pathlib import Path
|
|
14
|
+
from typing import Any, Literal, TypeAlias
|
|
15
|
+
|
|
16
|
+
import numpy as np
|
|
17
|
+
from dataeval.protocols import AnnotatedDataset, DatasetMetadata
|
|
18
|
+
from maite_datasets.adapters import HFImageClassificationDataset, HFObjectDetectionDataset
|
|
19
|
+
from numpy.typing import NDArray
|
|
20
|
+
from pydantic import BaseModel
|
|
21
|
+
|
|
22
|
+
logger: logging.Logger = logging.getLogger(__name__)
|
|
23
|
+
|
|
24
|
+
MaiteDataset: TypeAlias = HFImageClassificationDataset | HFObjectDetectionDataset
|
|
25
|
+
|
|
26
|
+
SUPPORTED_EXTENSIONS: frozenset[str] = frozenset({".png", ".jpg", ".jpeg", ".bmp", ".tif", ".tiff", ".webp"})
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class ImageFolderDataset:
|
|
30
|
+
"""MAITE-compatible dataset backed by a directory of image files.
|
|
31
|
+
|
|
32
|
+
Loads images lazily via PIL and converts to CHW numpy arrays. Returns
|
|
33
|
+
3-tuples ``(image, target, datum_metadata)``.
|
|
34
|
+
|
|
35
|
+
When ``infer_labels=False`` (default), target is an empty array (no
|
|
36
|
+
labels). When ``infer_labels=True``, immediate child directories of
|
|
37
|
+
``root`` are treated as class names and target is a one-hot float32
|
|
38
|
+
vector.
|
|
39
|
+
"""
|
|
40
|
+
|
|
41
|
+
def __init__(self, root: Path, *, recursive: bool = False, infer_labels: bool = False) -> None:
|
|
42
|
+
"""Initialize dataset from a directory of image files."""
|
|
43
|
+
self._root = root
|
|
44
|
+
if not root.is_dir():
|
|
45
|
+
raise FileNotFoundError(f"Image folder not found: {root}")
|
|
46
|
+
|
|
47
|
+
if infer_labels:
|
|
48
|
+
self._paths, self._labels, self._index2label = self._discover_labeled(root)
|
|
49
|
+
else:
|
|
50
|
+
self._paths = self._discover_unlabeled(root, recursive=recursive)
|
|
51
|
+
self._labels: list[int] | None = None
|
|
52
|
+
self._index2label: dict[int, str] = {}
|
|
53
|
+
|
|
54
|
+
if not self._paths:
|
|
55
|
+
raise FileNotFoundError(
|
|
56
|
+
f"No supported image files found in {root}. "
|
|
57
|
+
f"Supported extensions: {', '.join(sorted(SUPPORTED_EXTENSIONS))}"
|
|
58
|
+
)
|
|
59
|
+
logger.info("ImageFolderDataset: found %d images in %s", len(self._paths), root)
|
|
60
|
+
|
|
61
|
+
# -- Discovery --------------------------------------------------------
|
|
62
|
+
|
|
63
|
+
@staticmethod
|
|
64
|
+
def _discover_unlabeled(root: Path, *, recursive: bool) -> list[Path]:
|
|
65
|
+
"""Flat or recursive image discovery (no labels)."""
|
|
66
|
+
glob_pattern = "**/*" if recursive else "*"
|
|
67
|
+
return sorted(p for p in root.glob(glob_pattern) if p.is_file() and p.suffix.lower() in SUPPORTED_EXTENSIONS)
|
|
68
|
+
|
|
69
|
+
@staticmethod
|
|
70
|
+
def _discover_labeled(
|
|
71
|
+
root: Path,
|
|
72
|
+
) -> tuple[list[Path], list[int], dict[int, str]]:
|
|
73
|
+
"""Subdirectory-as-label discovery (torchvision ImageFolder convention).
|
|
74
|
+
|
|
75
|
+
Immediate child directories of *root* become class names, sorted
|
|
76
|
+
alphabetically and mapped to dense indices 0, 1, 2, …. Images in
|
|
77
|
+
each class directory are collected via ``rglob``. Empty class
|
|
78
|
+
directories are silently skipped.
|
|
79
|
+
"""
|
|
80
|
+
class_dirs = sorted(d for d in root.iterdir() if d.is_dir())
|
|
81
|
+
if not class_dirs:
|
|
82
|
+
raise FileNotFoundError(f"No class subdirectories found in {root}")
|
|
83
|
+
|
|
84
|
+
# Log top-level images that will be ignored
|
|
85
|
+
top_level_images = [p for p in root.iterdir() if p.is_file() and p.suffix.lower() in SUPPORTED_EXTENSIONS]
|
|
86
|
+
if top_level_images:
|
|
87
|
+
logger.debug(
|
|
88
|
+
"ImageFolderDataset: ignoring %d top-level images in labeled mode",
|
|
89
|
+
len(top_level_images),
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
paths: list[Path] = []
|
|
93
|
+
labels: list[int] = []
|
|
94
|
+
index2label: dict[int, str] = {}
|
|
95
|
+
class_idx = 0
|
|
96
|
+
for class_dir in class_dirs:
|
|
97
|
+
class_images = sorted(
|
|
98
|
+
p for p in class_dir.rglob("*") if p.is_file() and p.suffix.lower() in SUPPORTED_EXTENSIONS
|
|
99
|
+
)
|
|
100
|
+
if not class_images:
|
|
101
|
+
continue # skip empty class dirs
|
|
102
|
+
index2label[class_idx] = class_dir.name
|
|
103
|
+
paths.extend(class_images)
|
|
104
|
+
labels.extend([class_idx] * len(class_images))
|
|
105
|
+
class_idx += 1
|
|
106
|
+
|
|
107
|
+
return paths, labels, index2label
|
|
108
|
+
|
|
109
|
+
# -- AnnotatedDataset protocol ----------------------------------------
|
|
110
|
+
|
|
111
|
+
def __len__(self) -> int:
|
|
112
|
+
"""Return the number of images in the dataset."""
|
|
113
|
+
return len(self._paths)
|
|
114
|
+
|
|
115
|
+
def __getitem__(self, index: int) -> tuple[NDArray[Any], NDArray[Any], dict[str, Any]]:
|
|
116
|
+
"""Return (image, target, metadata) for the given index."""
|
|
117
|
+
if index < 0:
|
|
118
|
+
index += len(self._paths)
|
|
119
|
+
if index < 0 or index >= len(self._paths):
|
|
120
|
+
raise IndexError(f"Index {index} out of range for dataset of size {len(self)}")
|
|
121
|
+
img_array = self._load_image(index)
|
|
122
|
+
|
|
123
|
+
if self._labels is not None:
|
|
124
|
+
# Labeled mode — one-hot target
|
|
125
|
+
num_classes = len(self._index2label)
|
|
126
|
+
target: NDArray[Any] = np.zeros(num_classes, dtype=np.float32)
|
|
127
|
+
target[self._labels[index]] = 1.0
|
|
128
|
+
else:
|
|
129
|
+
# Unlabeled mode — empty target
|
|
130
|
+
target = np.empty(0, dtype=np.intp)
|
|
131
|
+
|
|
132
|
+
datum_metadata: dict[str, Any] = {
|
|
133
|
+
"id": index,
|
|
134
|
+
"filename": self._paths[index].name,
|
|
135
|
+
}
|
|
136
|
+
return img_array, target, datum_metadata
|
|
137
|
+
|
|
138
|
+
@property
|
|
139
|
+
def metadata(self) -> dict[str, Any]:
|
|
140
|
+
"""Dataset-level metadata (DatasetMetadata TypedDict shape)."""
|
|
141
|
+
return {
|
|
142
|
+
"id": 0,
|
|
143
|
+
"index2label": dict(self._index2label),
|
|
144
|
+
}
|
|
145
|
+
|
|
146
|
+
# -- Internal ---------------------------------------------------------
|
|
147
|
+
|
|
148
|
+
@lru_cache(maxsize=64) # noqa: B019
|
|
149
|
+
def _load_image(self, index: int) -> NDArray[Any]:
|
|
150
|
+
"""Load and convert a single image to CHW float32 numpy array."""
|
|
151
|
+
from PIL import Image
|
|
152
|
+
|
|
153
|
+
path = self._paths[index]
|
|
154
|
+
with Image.open(path) as img:
|
|
155
|
+
img = img.convert("RGB")
|
|
156
|
+
return np.transpose(
|
|
157
|
+
np.array(img, dtype=np.float32), # HWC → CHW
|
|
158
|
+
(2, 0, 1),
|
|
159
|
+
)
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
class _ObjectDetectionTarget:
|
|
163
|
+
"""Lightweight object-detection target conforming to :class:`~dataeval.protocols.ObjectDetectionTarget`."""
|
|
164
|
+
|
|
165
|
+
__slots__ = ("_boxes", "_labels", "_scores")
|
|
166
|
+
|
|
167
|
+
def __init__(self, boxes: NDArray[Any], labels: NDArray[Any], scores: NDArray[Any]) -> None:
|
|
168
|
+
self._boxes = boxes
|
|
169
|
+
self._labels = labels
|
|
170
|
+
self._scores = scores
|
|
171
|
+
|
|
172
|
+
@property
|
|
173
|
+
def boxes(self) -> NDArray[Any]:
|
|
174
|
+
""":class:`NDArray` of shape ``(N, 4)`` — XYXY bounding boxes."""
|
|
175
|
+
return self._boxes
|
|
176
|
+
|
|
177
|
+
@property
|
|
178
|
+
def labels(self) -> NDArray[Any]:
|
|
179
|
+
""":class:`NDArray` of shape ``(N,)`` — integer class labels."""
|
|
180
|
+
return self._labels
|
|
181
|
+
|
|
182
|
+
@property
|
|
183
|
+
def scores(self) -> NDArray[Any]:
|
|
184
|
+
""":class:`NDArray` of shape ``(N, M)`` — one-hot prediction scores."""
|
|
185
|
+
return self._scores
|
|
186
|
+
|
|
187
|
+
|
|
188
|
+
class TorchvisionDataset:
|
|
189
|
+
"""MAITE-compatible adapter for torchvision ``VisionDataset`` instances.
|
|
190
|
+
|
|
191
|
+
Wraps a torchvision dataset that returns ``(image, target)`` tuples and
|
|
192
|
+
converts them to the ``(NDArray, target, dict)`` format expected by the
|
|
193
|
+
:class:`~dataeval.protocols.AnnotatedDataset` protocol.
|
|
194
|
+
|
|
195
|
+
**Image classification** datasets return ``(image, int_label)``; the
|
|
196
|
+
integer label is converted to a one-hot float32 vector when *classes*
|
|
197
|
+
is discoverable on the wrapped dataset.
|
|
198
|
+
|
|
199
|
+
**Object detection** datasets (e.g. those wrapped with
|
|
200
|
+
``torchvision.datasets.wrap_dataset_for_transforms_v2``) return
|
|
201
|
+
``(image, dict)`` where the dict contains ``"boxes"``
|
|
202
|
+
(``BoundingBoxes``) and ``"labels"`` keys. Bounding boxes are
|
|
203
|
+
normalised to XYXY format regardless of the source
|
|
204
|
+
``BoundingBoxFormat``.
|
|
205
|
+
|
|
206
|
+
Parameters
|
|
207
|
+
----------
|
|
208
|
+
dataset : torch.utils.data.Dataset
|
|
209
|
+
A torchvision-style dataset whose ``__getitem__`` returns
|
|
210
|
+
``(image, target)`` tuples.
|
|
211
|
+
"""
|
|
212
|
+
|
|
213
|
+
def __init__(self, dataset: Any) -> None:
|
|
214
|
+
"""Initialize adapter from a torchvision dataset."""
|
|
215
|
+
self._dataset = dataset
|
|
216
|
+
|
|
217
|
+
# Discover class names from the dataset if available
|
|
218
|
+
classes: list[str] | None = getattr(dataset, "classes", None)
|
|
219
|
+
self._index2label: dict[int, str] = {i: c for i, c in enumerate(classes)} if classes else {}
|
|
220
|
+
self._num_classes: int | None = len(classes) if classes else None
|
|
221
|
+
|
|
222
|
+
name = getattr(dataset, "__class__", type(dataset)).__name__
|
|
223
|
+
logger.info("TorchvisionDataset: wrapping %s (%d samples)", name, len(self))
|
|
224
|
+
|
|
225
|
+
# -- AnnotatedDataset protocol ----------------------------------------
|
|
226
|
+
|
|
227
|
+
def __len__(self) -> int:
|
|
228
|
+
"""Return the number of samples in the wrapped dataset."""
|
|
229
|
+
return len(self._dataset) # type: ignore[arg-type]
|
|
230
|
+
|
|
231
|
+
def __getitem__(self, index: int) -> tuple[NDArray[Any], Any, dict[str, Any]]:
|
|
232
|
+
"""Return ``(image, target, metadata)`` for the given index."""
|
|
233
|
+
image, target = self._dataset[index]
|
|
234
|
+
|
|
235
|
+
img_array = self._convert_image(image)
|
|
236
|
+
|
|
237
|
+
if isinstance(target, dict) and "boxes" in target:
|
|
238
|
+
converted_target = self._convert_od_target(target)
|
|
239
|
+
else:
|
|
240
|
+
converted_target = self._convert_cls_target(target)
|
|
241
|
+
|
|
242
|
+
datum_metadata: dict[str, Any] = {"id": index}
|
|
243
|
+
return img_array, converted_target, datum_metadata
|
|
244
|
+
|
|
245
|
+
@property
|
|
246
|
+
def metadata(self) -> DatasetMetadata:
|
|
247
|
+
"""Dataset-level metadata (DatasetMetadata TypedDict shape)."""
|
|
248
|
+
name = getattr(self._dataset, "__class__", type(self._dataset)).__name__
|
|
249
|
+
return DatasetMetadata(
|
|
250
|
+
id=name,
|
|
251
|
+
index2label=dict(self._index2label),
|
|
252
|
+
)
|
|
253
|
+
|
|
254
|
+
# -- Internal ---------------------------------------------------------
|
|
255
|
+
|
|
256
|
+
@staticmethod
|
|
257
|
+
def _convert_image(image: Any) -> NDArray[Any]:
|
|
258
|
+
"""Convert a PIL Image or torch Tensor to CHW float32 numpy array."""
|
|
259
|
+
from PIL import Image
|
|
260
|
+
|
|
261
|
+
if isinstance(image, Image.Image):
|
|
262
|
+
image = image.convert("RGB")
|
|
263
|
+
return np.transpose(np.array(image, dtype=np.float32), (2, 0, 1))
|
|
264
|
+
|
|
265
|
+
# torch.Tensor — detach, move to cpu, convert
|
|
266
|
+
arr = np.asarray(image, dtype=np.float32) # handles Tensor and ndarray
|
|
267
|
+
if arr.ndim == 3 and arr.shape[2] in (1, 3, 4):
|
|
268
|
+
# HWC → CHW
|
|
269
|
+
arr = np.transpose(arr, (2, 0, 1))
|
|
270
|
+
return arr
|
|
271
|
+
|
|
272
|
+
def _convert_cls_target(self, target: Any) -> NDArray[Any]:
|
|
273
|
+
"""Convert an integer label to a one-hot vector, or pass through arrays."""
|
|
274
|
+
if isinstance(target, int) and self._num_classes is not None:
|
|
275
|
+
one_hot = np.zeros(self._num_classes, dtype=np.float32)
|
|
276
|
+
one_hot[target] = 1.0
|
|
277
|
+
return one_hot
|
|
278
|
+
return np.asarray(target, dtype=np.float32)
|
|
279
|
+
|
|
280
|
+
def _convert_od_target(self, target: dict[str, Any]) -> _ObjectDetectionTarget:
|
|
281
|
+
"""Convert a torchvision v2 object-detection target dict.
|
|
282
|
+
|
|
283
|
+
Expects *target* to contain at least ``"boxes"`` and ``"labels"``
|
|
284
|
+
keys. ``"boxes"`` may be a ``torchvision.tv_tensors.BoundingBoxes``
|
|
285
|
+
(with an associated format) or a plain tensor already in XYXY order.
|
|
286
|
+
All non-XYXY formats are converted to XYXY.
|
|
287
|
+
"""
|
|
288
|
+
boxes_raw = target["boxes"]
|
|
289
|
+
labels_raw = target["labels"]
|
|
290
|
+
|
|
291
|
+
# Convert boxes to XYXY numpy — respect BoundingBoxes.format
|
|
292
|
+
boxes_np = self._boxes_to_xyxy_numpy(boxes_raw)
|
|
293
|
+
labels_np = np.asarray(labels_raw, dtype=np.intp)
|
|
294
|
+
|
|
295
|
+
# Build one-hot score matrix
|
|
296
|
+
num_classes = self._num_classes or (int(labels_np.max()) + 1 if len(labels_np) > 0 else 0)
|
|
297
|
+
scores_np = np.zeros((len(labels_np), num_classes), dtype=np.float32)
|
|
298
|
+
for i, lbl in enumerate(labels_np):
|
|
299
|
+
scores_np[i, lbl] = 1.0
|
|
300
|
+
|
|
301
|
+
return _ObjectDetectionTarget(boxes=boxes_np, labels=labels_np, scores=scores_np)
|
|
302
|
+
|
|
303
|
+
@staticmethod
|
|
304
|
+
def _boxes_to_xyxy_numpy(boxes: Any) -> NDArray[np.float32]:
|
|
305
|
+
"""Convert bounding boxes to XYXY float32 numpy array.
|
|
306
|
+
|
|
307
|
+
Handles ``torchvision.tv_tensors.BoundingBoxes`` (any format),
|
|
308
|
+
plain torch Tensors (assumed XYXY), and numpy arrays.
|
|
309
|
+
"""
|
|
310
|
+
try:
|
|
311
|
+
from torchvision.tv_tensors import BoundingBoxes, BoundingBoxFormat
|
|
312
|
+
|
|
313
|
+
if isinstance(boxes, BoundingBoxes) and boxes.format != BoundingBoxFormat.XYXY:
|
|
314
|
+
from torchvision.ops import box_convert
|
|
315
|
+
|
|
316
|
+
boxes = box_convert(boxes, in_fmt=boxes.format.name.lower(), out_fmt="xyxy")
|
|
317
|
+
except ImportError: # pragma: no cover — torchvision not installed
|
|
318
|
+
pass
|
|
319
|
+
|
|
320
|
+
return np.asarray(boxes, dtype=np.float32).reshape(-1, 4)
|
|
321
|
+
|
|
322
|
+
|
|
323
|
+
def load_dataset_torchvision(dataset: Any) -> TorchvisionDataset:
|
|
324
|
+
"""Wrap a torchvision dataset as a MAITE-compatible dataset.
|
|
325
|
+
|
|
326
|
+
Parameters
|
|
327
|
+
----------
|
|
328
|
+
dataset : torch.utils.data.Dataset
|
|
329
|
+
A torchvision-style dataset.
|
|
330
|
+
|
|
331
|
+
Returns
|
|
332
|
+
-------
|
|
333
|
+
TorchvisionDataset
|
|
334
|
+
MAITE-compatible wrapper around the torchvision dataset.
|
|
335
|
+
"""
|
|
336
|
+
return TorchvisionDataset(dataset)
|
|
337
|
+
|
|
338
|
+
|
|
339
|
+
def load_dataset_huggingface(path: Path, split: str | None = None) -> MaiteDataset:
|
|
340
|
+
"""Load a HuggingFace dataset and convert to MAITE format.
|
|
341
|
+
|
|
342
|
+
Parameters
|
|
343
|
+
----------
|
|
344
|
+
path : Path
|
|
345
|
+
Path to the dataset directory containing HuggingFace dataset files.
|
|
346
|
+
|
|
347
|
+
Returns
|
|
348
|
+
-------
|
|
349
|
+
MaiteDataset
|
|
350
|
+
MAITE-compatible dataset object that can be used with DataEval.
|
|
351
|
+
|
|
352
|
+
Raises
|
|
353
|
+
------
|
|
354
|
+
ImportError
|
|
355
|
+
If required dependencies (datasets, maite_datasets) are missing.
|
|
356
|
+
RuntimeError
|
|
357
|
+
If dataset loading or conversion fails.
|
|
358
|
+
KeyError
|
|
359
|
+
If split is not specified for a multi-split dataset or the specified split is not found in the dataset.
|
|
360
|
+
|
|
361
|
+
Examples
|
|
362
|
+
--------
|
|
363
|
+
>>> from pathlib import Path
|
|
364
|
+
>>> from dataeval_flow import load_dataset_huggingface
|
|
365
|
+
>>> ds = load_dataset_huggingface(Path("/data/cifar10"))
|
|
366
|
+
"""
|
|
367
|
+
from datasets import load_from_disk
|
|
368
|
+
from maite_datasets.adapters import from_huggingface
|
|
369
|
+
|
|
370
|
+
dataset = load_from_disk(str(path))
|
|
371
|
+
|
|
372
|
+
logger.info("Loaded type: %s", type(dataset).__name__)
|
|
373
|
+
# union-attr: load_from_disk returns Dataset | DatasetDict; .keys() only on DatasetDict.
|
|
374
|
+
if hasattr(dataset, "keys") and callable(dataset.keys): # type: ignore[union-attr]
|
|
375
|
+
available_splits = list(dataset.keys()) # type: ignore[union-attr]
|
|
376
|
+
if split is None or split not in available_splits:
|
|
377
|
+
raise KeyError(f"Requested split '{split}' not found in dataset. Available splits: {available_splits}")
|
|
378
|
+
logger.info("DatasetDict detected with %d splits: %s", len(available_splits), available_splits)
|
|
379
|
+
logger.info("Selecting split '%s' as specified in config.", split)
|
|
380
|
+
dataset = dataset[split]
|
|
381
|
+
else:
|
|
382
|
+
logger.info("Dataset detected (single split)")
|
|
383
|
+
|
|
384
|
+
# load_from_disk returns Dataset | DatasetDict; after the dict guard above
|
|
385
|
+
# it's a single Dataset, but pyright can't narrow through hasattr.
|
|
386
|
+
return from_huggingface(dataset) # type: ignore[arg-type]
|
|
387
|
+
|
|
388
|
+
|
|
389
|
+
def load_dataset_image_folder(path: Path, *, recursive: bool = False, infer_labels: bool = False) -> ImageFolderDataset:
|
|
390
|
+
"""Load an image folder dataset."""
|
|
391
|
+
return ImageFolderDataset(path, recursive=recursive, infer_labels=infer_labels)
|
|
392
|
+
|
|
393
|
+
|
|
394
|
+
def load_dataset_coco(
|
|
395
|
+
path: Path,
|
|
396
|
+
*,
|
|
397
|
+
annotations_file: str | None = None,
|
|
398
|
+
images_dir: str | None = None,
|
|
399
|
+
classes_file: str | None = None,
|
|
400
|
+
) -> Any:
|
|
401
|
+
"""Load a COCO-format object detection dataset.
|
|
402
|
+
|
|
403
|
+
Parameters
|
|
404
|
+
----------
|
|
405
|
+
path : Path
|
|
406
|
+
Root directory of the COCO dataset.
|
|
407
|
+
annotations_file : str | None
|
|
408
|
+
Name of the annotations JSON file (default: reader's default).
|
|
409
|
+
images_dir : str | None
|
|
410
|
+
Name of the images subdirectory (default: reader's default).
|
|
411
|
+
classes_file : str | None
|
|
412
|
+
Name of the classes text file (default: reader's default).
|
|
413
|
+
|
|
414
|
+
Returns
|
|
415
|
+
-------
|
|
416
|
+
Any
|
|
417
|
+
MAITE-compatible object detection dataset (``COCODataset``).
|
|
418
|
+
"""
|
|
419
|
+
from maite_datasets.object_detection import COCODatasetReader
|
|
420
|
+
|
|
421
|
+
kwargs: dict[str, str] = {}
|
|
422
|
+
if annotations_file is not None:
|
|
423
|
+
kwargs["annotation_file"] = annotations_file # reader uses singular
|
|
424
|
+
if images_dir is not None:
|
|
425
|
+
kwargs["images_dir"] = images_dir
|
|
426
|
+
if classes_file is not None:
|
|
427
|
+
kwargs["classes_file"] = classes_file
|
|
428
|
+
reader = COCODatasetReader(path, **kwargs)
|
|
429
|
+
dataset = reader.create_dataset()
|
|
430
|
+
logger.info("COCODataset: loaded %d images from %s", len(dataset), path)
|
|
431
|
+
return dataset
|
|
432
|
+
|
|
433
|
+
|
|
434
|
+
def load_dataset_yolo(
|
|
435
|
+
path: Path,
|
|
436
|
+
*,
|
|
437
|
+
images_dir: str | None = None,
|
|
438
|
+
labels_dir: str | None = None,
|
|
439
|
+
classes_file: str | None = None,
|
|
440
|
+
) -> Any:
|
|
441
|
+
"""Load a YOLO-format object detection dataset.
|
|
442
|
+
|
|
443
|
+
Parameters
|
|
444
|
+
----------
|
|
445
|
+
path : Path
|
|
446
|
+
Root directory of the YOLO dataset.
|
|
447
|
+
images_dir : str | None
|
|
448
|
+
Name of the images subdirectory (default: reader's default).
|
|
449
|
+
labels_dir : str | None
|
|
450
|
+
Name of the labels subdirectory (default: reader's default).
|
|
451
|
+
classes_file : str | None
|
|
452
|
+
Name of the classes text file (default: reader's default).
|
|
453
|
+
|
|
454
|
+
Returns
|
|
455
|
+
-------
|
|
456
|
+
Any
|
|
457
|
+
MAITE-compatible object detection dataset (``YOLODataset``).
|
|
458
|
+
"""
|
|
459
|
+
from maite_datasets.object_detection import YOLODatasetReader
|
|
460
|
+
|
|
461
|
+
kwargs: dict[str, str] = {}
|
|
462
|
+
if images_dir is not None:
|
|
463
|
+
kwargs["images_dir"] = images_dir
|
|
464
|
+
if labels_dir is not None:
|
|
465
|
+
kwargs["labels_dir"] = labels_dir
|
|
466
|
+
if classes_file is not None:
|
|
467
|
+
kwargs["classes_file"] = classes_file
|
|
468
|
+
reader = YOLODatasetReader(path, **kwargs) # type: ignore[arg-type] # kwargs are all str; pyright flags image_extensions
|
|
469
|
+
dataset = reader.create_dataset()
|
|
470
|
+
logger.info("YOLODataset: loaded %d images from %s", len(dataset), path)
|
|
471
|
+
return dataset
|
|
472
|
+
|
|
473
|
+
|
|
474
|
+
def load_dataset(
|
|
475
|
+
path: Path,
|
|
476
|
+
split: str | None = None,
|
|
477
|
+
dataset_format: Literal["huggingface", "coco", "yolo", "image_folder"] = "huggingface",
|
|
478
|
+
*,
|
|
479
|
+
recursive: bool = False,
|
|
480
|
+
infer_labels: bool = False,
|
|
481
|
+
annotations_file: str | None = None,
|
|
482
|
+
images_dir: str | None = None,
|
|
483
|
+
labels_dir: str | None = None,
|
|
484
|
+
classes_file: str | None = None,
|
|
485
|
+
) -> Any:
|
|
486
|
+
"""Load a dataset and convert to MAITE format.
|
|
487
|
+
|
|
488
|
+
This is the main entry point for loading datasets. Dispatches to the
|
|
489
|
+
appropriate loader based on ``dataset_format``.
|
|
490
|
+
|
|
491
|
+
Parameters
|
|
492
|
+
----------
|
|
493
|
+
path : Path
|
|
494
|
+
Path to the dataset directory.
|
|
495
|
+
split : str | None
|
|
496
|
+
Optional split name to load (e.g. "train", "test").
|
|
497
|
+
dataset_format : Literal["huggingface", "coco", "yolo", "image_folder"]
|
|
498
|
+
Dataset format identifier (default ``"huggingface"``).
|
|
499
|
+
recursive : bool
|
|
500
|
+
Scan subdirectories for images (image_folder only).
|
|
501
|
+
infer_labels : bool
|
|
502
|
+
Treat subdirectories as class labels (image_folder only).
|
|
503
|
+
annotations_file : str | None
|
|
504
|
+
Annotations file name (COCO only).
|
|
505
|
+
images_dir : str | None
|
|
506
|
+
Images subdirectory name (COCO/YOLO only).
|
|
507
|
+
labels_dir : str | None
|
|
508
|
+
Labels subdirectory name (YOLO only).
|
|
509
|
+
classes_file : str | None
|
|
510
|
+
Classes file name (COCO/YOLO only).
|
|
511
|
+
|
|
512
|
+
Returns
|
|
513
|
+
-------
|
|
514
|
+
MaiteDataset | ImageFolderDataset | Any
|
|
515
|
+
MAITE-compatible dataset object.
|
|
516
|
+
|
|
517
|
+
Raises
|
|
518
|
+
------
|
|
519
|
+
KeyError
|
|
520
|
+
If split is not specified for a multi-split dataset or the specified split is not found in the dataset.
|
|
521
|
+
ValueError
|
|
522
|
+
If the dataset format is not supported.
|
|
523
|
+
|
|
524
|
+
Examples
|
|
525
|
+
--------
|
|
526
|
+
>>> from pathlib import Path
|
|
527
|
+
>>> from dataeval_flow import load_dataset
|
|
528
|
+
>>> ds = load_dataset(Path("/data/cifar10"))
|
|
529
|
+
"""
|
|
530
|
+
if dataset_format == "huggingface":
|
|
531
|
+
return load_dataset_huggingface(path, split=split)
|
|
532
|
+
if dataset_format == "image_folder":
|
|
533
|
+
return load_dataset_image_folder(path, recursive=recursive, infer_labels=infer_labels)
|
|
534
|
+
if dataset_format == "coco":
|
|
535
|
+
return load_dataset_coco(
|
|
536
|
+
path, annotations_file=annotations_file, images_dir=images_dir, classes_file=classes_file
|
|
537
|
+
)
|
|
538
|
+
if dataset_format == "yolo":
|
|
539
|
+
return load_dataset_yolo(path, images_dir=images_dir, labels_dir=labels_dir, classes_file=classes_file)
|
|
540
|
+
msg = f"Unsupported dataset format: {dataset_format!r}"
|
|
541
|
+
raise ValueError(msg)
|
|
542
|
+
|
|
543
|
+
|
|
544
|
+
# ---------------------------------------------------------------------------
|
|
545
|
+
# Resolved dataset — unified output of config → dataset resolution
|
|
546
|
+
# ---------------------------------------------------------------------------
|
|
547
|
+
|
|
548
|
+
_LABEL_SOURCE: dict[str, str] = {
|
|
549
|
+
"coco": "annotations",
|
|
550
|
+
"yolo": "annotations",
|
|
551
|
+
"huggingface": "huggingface",
|
|
552
|
+
"maite": "protocol",
|
|
553
|
+
"torchvision": "torchvision",
|
|
554
|
+
}
|
|
555
|
+
|
|
556
|
+
|
|
557
|
+
@dataclass
|
|
558
|
+
class ResolvedDataset:
|
|
559
|
+
"""Result of resolving any dataset config into a ready-to-use dataset.
|
|
560
|
+
|
|
561
|
+
Produced by :func:`resolve_dataset` so that downstream code (orchestrator,
|
|
562
|
+
cache) never needs to branch on config type.
|
|
563
|
+
"""
|
|
564
|
+
|
|
565
|
+
name: str
|
|
566
|
+
dataset: AnnotatedDataset[Any]
|
|
567
|
+
label_source: str | None
|
|
568
|
+
cache_key: str
|
|
569
|
+
|
|
570
|
+
|
|
571
|
+
def resolve_dataset(config: BaseModel, data_dir: Path | None = None) -> ResolvedDataset:
|
|
572
|
+
"""Resolve a dataset config into a :class:`ResolvedDataset`.
|
|
573
|
+
|
|
574
|
+
Handles both file-backed (:class:`DatasetConfig` union members) and
|
|
575
|
+
in-memory (:class:`DatasetProtocolConfig`) configs, centralizing all
|
|
576
|
+
format-specific branching in one place.
|
|
577
|
+
|
|
578
|
+
Parameters
|
|
579
|
+
----------
|
|
580
|
+
config : BaseModel
|
|
581
|
+
Dataset configuration object.
|
|
582
|
+
data_dir : Path | None
|
|
583
|
+
Root directory for resolving relative dataset paths.
|
|
584
|
+
"""
|
|
585
|
+
from dataeval_flow.cache import dataset_fingerprint
|
|
586
|
+
from dataeval_flow.config.schemas._dataset import (
|
|
587
|
+
DatasetProtocolConfig,
|
|
588
|
+
HuggingFaceDatasetConfig,
|
|
589
|
+
ImageFolderDatasetConfig,
|
|
590
|
+
_DatasetConfigBase,
|
|
591
|
+
)
|
|
592
|
+
|
|
593
|
+
if isinstance(config, DatasetProtocolConfig):
|
|
594
|
+
dataset = load_dataset_torchvision(config.dataset) if config.format == "torchvision" else config.dataset
|
|
595
|
+
label_source: str | None = _LABEL_SOURCE.get(config.format)
|
|
596
|
+
cache_key = f"{config.name}:{config.format}:{config.version}"
|
|
597
|
+
elif isinstance(config, _DatasetConfigBase):
|
|
598
|
+
# All file-backed dataset configs share path/format; dispatch through load_dataset
|
|
599
|
+
kwargs: dict[str, Any] = {}
|
|
600
|
+
if isinstance(config, HuggingFaceDatasetConfig):
|
|
601
|
+
kwargs["split"] = config.split
|
|
602
|
+
elif isinstance(config, ImageFolderDatasetConfig):
|
|
603
|
+
kwargs["recursive"] = config.recursive
|
|
604
|
+
kwargs["infer_labels"] = config.infer_labels
|
|
605
|
+
else:
|
|
606
|
+
# Coco / Yolo — forward format-specific fields
|
|
607
|
+
for field_name in ("annotations_file", "images_dir", "labels_dir", "classes_file"):
|
|
608
|
+
if hasattr(config, field_name):
|
|
609
|
+
kwargs[field_name] = getattr(config, field_name)
|
|
610
|
+
|
|
611
|
+
from dataeval_flow.config._loader import resolve_path
|
|
612
|
+
|
|
613
|
+
dataset_path = resolve_path(config.path, data_dir)
|
|
614
|
+
|
|
615
|
+
dataset = load_dataset(dataset_path, dataset_format=config.format, **kwargs)
|
|
616
|
+
|
|
617
|
+
if isinstance(config, ImageFolderDatasetConfig) and config.infer_labels:
|
|
618
|
+
label_source = "filepath"
|
|
619
|
+
else:
|
|
620
|
+
label_source = _LABEL_SOURCE.get(config.format)
|
|
621
|
+
cache_key = config.model_dump_json(exclude_defaults=False)
|
|
622
|
+
else:
|
|
623
|
+
raise ValueError(f"Unsupported dataset config type: {type(config).__name__}")
|
|
624
|
+
|
|
625
|
+
# Append a content-based fingerprint so the cache invalidates when
|
|
626
|
+
# the underlying data changes even if the config metadata is unchanged.
|
|
627
|
+
fingerprint = dataset_fingerprint(dataset)
|
|
628
|
+
cache_key = f"{cache_key}|fp:{fingerprint}"
|
|
629
|
+
|
|
630
|
+
return ResolvedDataset(
|
|
631
|
+
name=config.name,
|
|
632
|
+
dataset=dataset,
|
|
633
|
+
label_source=label_source,
|
|
634
|
+
cache_key=cache_key,
|
|
635
|
+
)
|