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.
Files changed (94) hide show
  1. dataeval_flow/__init__.py +93 -0
  2. dataeval_flow/__main__.py +149 -0
  3. dataeval_flow/_app/__init__.py +5 -0
  4. dataeval_flow/_app/_model/__init__.py +5 -0
  5. dataeval_flow/_app/_model/_coerce.py +126 -0
  6. dataeval_flow/_app/_model/_discover.py +171 -0
  7. dataeval_flow/_app/_model/_execution.py +108 -0
  8. dataeval_flow/_app/_model/_introspect.py +280 -0
  9. dataeval_flow/_app/_model/_item.py +213 -0
  10. dataeval_flow/_app/_model/_registry.py +255 -0
  11. dataeval_flow/_app/_model/_state.py +322 -0
  12. dataeval_flow/_app/_model/_undo.py +61 -0
  13. dataeval_flow/_app/_panes/__init__.py +35 -0
  14. dataeval_flow/_app/_panes/_config_pane.py +173 -0
  15. dataeval_flow/_app/_panes/_result_pane.py +125 -0
  16. dataeval_flow/_app/_panes/_task_pane.py +91 -0
  17. dataeval_flow/_app/_panes/_widgets.py +111 -0
  18. dataeval_flow/_app/_screens/__init__.py +25 -0
  19. dataeval_flow/_app/_screens/_base.py +242 -0
  20. dataeval_flow/_app/_screens/_detail.py +333 -0
  21. dataeval_flow/_app/_screens/_model.py +102 -0
  22. dataeval_flow/_app/_screens/_params.py +80 -0
  23. dataeval_flow/_app/_screens/_pathpicker.py +68 -0
  24. dataeval_flow/_app/_screens/_section.py +621 -0
  25. dataeval_flow/_app/_screens/_settings.py +183 -0
  26. dataeval_flow/_app/_viewmodel/__init__.py +15 -0
  27. dataeval_flow/_app/_viewmodel/_builder_vm.py +272 -0
  28. dataeval_flow/_app/_viewmodel/_model_vm.py +70 -0
  29. dataeval_flow/_app/_viewmodel/_rendering.py +189 -0
  30. dataeval_flow/_app/_viewmodel/_result_vm.py +210 -0
  31. dataeval_flow/_app/_viewmodel/_section_vm.py +224 -0
  32. dataeval_flow/_app/app.py +742 -0
  33. dataeval_flow/_app/cli.py +592 -0
  34. dataeval_flow/_logging.py +102 -0
  35. dataeval_flow/cache.py +1355 -0
  36. dataeval_flow/config/__init__.py +80 -0
  37. dataeval_flow/config/_loader.py +79 -0
  38. dataeval_flow/config/_merge.py +92 -0
  39. dataeval_flow/config/_models.py +115 -0
  40. dataeval_flow/config/_paths.py +85 -0
  41. dataeval_flow/config/schemas/__init__.py +112 -0
  42. dataeval_flow/config/schemas/_dataset.py +111 -0
  43. dataeval_flow/config/schemas/_extractor.py +119 -0
  44. dataeval_flow/config/schemas/_metadata.py +28 -0
  45. dataeval_flow/config/schemas/_preprocessor.py +18 -0
  46. dataeval_flow/config/schemas/_selection.py +100 -0
  47. dataeval_flow/config/schemas/_task.py +89 -0
  48. dataeval_flow/config/schemas/_workflow.py +135 -0
  49. dataeval_flow/dataset.py +635 -0
  50. dataeval_flow/embeddings.py +135 -0
  51. dataeval_flow/metadata.py +48 -0
  52. dataeval_flow/preprocessing.py +141 -0
  53. dataeval_flow/py.typed +0 -0
  54. dataeval_flow/runner.py +118 -0
  55. dataeval_flow/selection.py +50 -0
  56. dataeval_flow/workflow/__init__.py +328 -0
  57. dataeval_flow/workflow/_text_report.py +511 -0
  58. dataeval_flow/workflow/base.py +69 -0
  59. dataeval_flow/workflow/orchestrator.py +454 -0
  60. dataeval_flow/workflows/__init__.py +1 -0
  61. dataeval_flow/workflows/analysis/__init__.py +38 -0
  62. dataeval_flow/workflows/analysis/outputs.py +202 -0
  63. dataeval_flow/workflows/analysis/params.py +114 -0
  64. dataeval_flow/workflows/analysis/workflow.py +1313 -0
  65. dataeval_flow/workflows/cleaning/__init__.py +23 -0
  66. dataeval_flow/workflows/cleaning/outputs.py +200 -0
  67. dataeval_flow/workflows/cleaning/params.py +160 -0
  68. dataeval_flow/workflows/cleaning/report.py +304 -0
  69. dataeval_flow/workflows/cleaning/workflow.py +794 -0
  70. dataeval_flow/workflows/drift/__init__.py +1 -0
  71. dataeval_flow/workflows/drift/outputs.py +144 -0
  72. dataeval_flow/workflows/drift/params.py +332 -0
  73. dataeval_flow/workflows/drift/report.py +201 -0
  74. dataeval_flow/workflows/drift/workflow.py +647 -0
  75. dataeval_flow/workflows/ood/__init__.py +1 -0
  76. dataeval_flow/workflows/ood/outputs.py +134 -0
  77. dataeval_flow/workflows/ood/params.py +161 -0
  78. dataeval_flow/workflows/ood/report.py +311 -0
  79. dataeval_flow/workflows/ood/workflow.py +728 -0
  80. dataeval_flow/workflows/prioritization/__init__.py +1 -0
  81. dataeval_flow/workflows/prioritization/outputs.py +122 -0
  82. dataeval_flow/workflows/prioritization/params.py +124 -0
  83. dataeval_flow/workflows/prioritization/report.py +117 -0
  84. dataeval_flow/workflows/prioritization/workflow.py +587 -0
  85. dataeval_flow/workflows/splitting/__init__.py +25 -0
  86. dataeval_flow/workflows/splitting/outputs.py +101 -0
  87. dataeval_flow/workflows/splitting/params.py +61 -0
  88. dataeval_flow/workflows/splitting/report.py +485 -0
  89. dataeval_flow/workflows/splitting/workflow.py +371 -0
  90. dataeval_flow-0.1.0.dist-info/METADATA +305 -0
  91. dataeval_flow-0.1.0.dist-info/RECORD +94 -0
  92. dataeval_flow-0.1.0.dist-info/WHEEL +4 -0
  93. dataeval_flow-0.1.0.dist-info/entry_points.txt +2 -0
  94. dataeval_flow-0.1.0.dist-info/licenses/LICENSE +21 -0
@@ -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
+ )