genml-kit 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.
- genml_kit/__init__.py +11 -0
- genml_kit/datasets/__init__.py +11 -0
- genml_kit/datasets/balanced_sampler.py +82 -0
- genml_kit/datasets/ensemble.py +355 -0
- genml_kit/datasets/field_dataset.py +40 -0
- genml_kit/datasets/hf_proxy.py +239 -0
- genml_kit/datasets/image_folder.py +264 -0
- genml_kit/datasets/retry.py +43 -0
- genml_kit/datasets/weighted_sampler.py +70 -0
- genml_kit/io/__init__.py +1 -0
- genml_kit/io/checkpointing.py +721 -0
- genml_kit/io/storage_utils.py +339 -0
- genml_kit/models/__init__.py +37 -0
- genml_kit/models/attention_pooling.py +43 -0
- genml_kit/models/byol.py +98 -0
- genml_kit/models/cls_model_wrapper/__init__.py +9 -0
- genml_kit/models/cls_model_wrapper/loader.py +50 -0
- genml_kit/models/cls_model_wrapper/model.py +102 -0
- genml_kit/models/contrastive.py +74 -0
- genml_kit/models/convvit/__init__.py +13 -0
- genml_kit/models/convvit/loader.py +162 -0
- genml_kit/models/convvit/masked_encoder.py +32 -0
- genml_kit/models/convvit/model.py +209 -0
- genml_kit/models/convvit/processor.py +47 -0
- genml_kit/models/dino.py +151 -0
- genml_kit/models/encoder_utils.py +83 -0
- genml_kit/models/masked_encoder.py +45 -0
- genml_kit/models/processors/__init__.py +5 -0
- genml_kit/models/processors/base.py +76 -0
- genml_kit/models/registry.py +231 -0
- genml_kit/models/simmim.py +122 -0
- genml_kit/models/timm/__init__.py +3 -0
- genml_kit/models/timm/loader.py +110 -0
- genml_kit/models/timm/model.py +59 -0
- genml_kit/models/timm/processor.py +61 -0
- genml_kit/models/transformer_utils.py +86 -0
- genml_kit/models/uvito/__init__.py +11 -0
- genml_kit/models/uvito/loader.py +116 -0
- genml_kit/models/uvito/model.py +150 -0
- genml_kit/models/uvito/processor.py +33 -0
- genml_kit/pretrain/__init__.py +1 -0
- genml_kit/pretrain/augmentations/__init__.py +1 -0
- genml_kit/pretrain/augmentations/dual_view.py +11 -0
- genml_kit/pretrain/augmentations/multicrop.py +105 -0
- genml_kit/pretrain/cli.py +884 -0
- genml_kit/pretrain/losses/__init__.py +8 -0
- genml_kit/pretrain/losses/byol.py +24 -0
- genml_kit/pretrain/losses/contrastive.py +54 -0
- genml_kit/pretrain/losses/dino.py +82 -0
- genml_kit/pretrain/losses/focal.py +77 -0
- genml_kit/pretrain/methods/__init__.py +11 -0
- genml_kit/pretrain/methods/base.py +100 -0
- genml_kit/pretrain/methods/byol.py +105 -0
- genml_kit/pretrain/methods/dino.py +179 -0
- genml_kit/pretrain/methods/ijepa.py +310 -0
- genml_kit/pretrain/methods/registry.py +33 -0
- genml_kit/pretrain/methods/simmim.py +138 -0
- genml_kit/pretrain/methods/supcon.py +72 -0
- genml_kit/training/__init__.py +1 -0
- genml_kit/training/classifiers/__init__.py +80 -0
- genml_kit/training/classifiers/base.py +45 -0
- genml_kit/training/classifiers/cls_attention.py +91 -0
- genml_kit/training/classifiers/mlp.py +54 -0
- genml_kit/training/eval.py +128 -0
- genml_kit/training/grad_monitor.py +434 -0
- genml_kit/training/infer.py +222 -0
- genml_kit/training/metrics.py +76 -0
- genml_kit/training/model_utils.py +456 -0
- genml_kit/training/optim_factory.py +459 -0
- genml_kit/training/param_align.py +166 -0
- genml_kit/training/train.py +1381 -0
- genml_kit/training/train_reporting.py +197 -0
- genml_kit/training/tta.py +91 -0
- genml_kit/training/xgb_pipeline.py +153 -0
- genml_kit/training/xgb_utils.py +144 -0
- genml_kit/utils/__init__.py +1 -0
- genml_kit/utils/args.py +176 -0
- genml_kit/utils/attr.py +141 -0
- genml_kit/utils/cli.py +108 -0
- genml_kit/utils/gpu.py +41 -0
- genml_kit/utils/image_dump.py +110 -0
- genml_kit/utils/label.py +27 -0
- genml_kit/utils/logging.py +150 -0
- genml_kit/utils/script.py +82 -0
- genml_kit/utils/seed.py +86 -0
- genml_kit/utils/signal.py +114 -0
- genml_kit/utils/table.py +58 -0
- genml_kit/utils/transformer.py +57 -0
- genml_kit-0.1.0.dist-info/METADATA +1209 -0
- genml_kit-0.1.0.dist-info/RECORD +94 -0
- genml_kit-0.1.0.dist-info/WHEEL +5 -0
- genml_kit-0.1.0.dist-info/entry_points.txt +4 -0
- genml_kit-0.1.0.dist-info/licenses/LICENSE +190 -0
- genml_kit-0.1.0.dist-info/top_level.txt +1 -0
genml_kit/__init__.py
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
1
|
+
"""General-purpose ML toolkit.
|
|
2
|
+
|
|
3
|
+
``genml_kit`` provides the model-agnostic core of the toolkit: training
|
|
4
|
+
and inference harnesses, self-supervised pre-training methods, dataset
|
|
5
|
+
utilities, model registries, checkpointing, and storage helpers.
|
|
6
|
+
|
|
7
|
+
Domain-specific helpers (e.g. skin-lesion dataset preparation) live in
|
|
8
|
+
the ``genml_kit/`` folder of the repository.
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
__version__ = "0.1.0"
|
|
@@ -0,0 +1,11 @@
|
|
|
1
|
+
"""Dataset utilities.
|
|
2
|
+
|
|
3
|
+
Provides :class:`DatasetEnsemble` for stitching together multiple image
|
|
4
|
+
datasets into a unified pre-training corpus, and :class:`HFDatasetProxy`
|
|
5
|
+
for bridging HuggingFace datasets to PyTorch.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from genml_kit.datasets.ensemble import DatasetEnsemble
|
|
9
|
+
from genml_kit.datasets.hf_proxy import HFDatasetProxy
|
|
10
|
+
|
|
11
|
+
__all__ = ["DatasetEnsemble", "HFDatasetProxy"]
|
|
@@ -0,0 +1,82 @@
|
|
|
1
|
+
"""BalancedBatchSampler — ensures uniform class representation per batch."""
|
|
2
|
+
|
|
3
|
+
import logging
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
from torch.utils.data import Sampler
|
|
7
|
+
|
|
8
|
+
from genml_kit.utils.logging import fatal
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class BalancedBatchSampler(Sampler):
|
|
12
|
+
"""Yields mini-batches with a guaranteed number of samples per class.
|
|
13
|
+
|
|
14
|
+
Each batch contains ``batch_size // samples_per_class`` groups, where
|
|
15
|
+
every group holds ``samples_per_class`` examples from the same class.
|
|
16
|
+
Classes are sampled uniformly at random so that rare classes appear
|
|
17
|
+
with the same frequency as common ones.
|
|
18
|
+
|
|
19
|
+
Parameters
|
|
20
|
+
----------
|
|
21
|
+
labels : array-like of int
|
|
22
|
+
Integer class label for every sample in the dataset.
|
|
23
|
+
batch_size : int
|
|
24
|
+
Total batch size (must be divisible by *samples_per_class*).
|
|
25
|
+
samples_per_class : int
|
|
26
|
+
Number of examples drawn from each class within a batch.
|
|
27
|
+
seed : int, optional
|
|
28
|
+
Seed for the class/index RNG. When None (default) each epoch's
|
|
29
|
+
batch composition is drawn from fresh OS entropy; pass a seed for
|
|
30
|
+
reproducible batches (a fresh generator is created per ``__iter__``,
|
|
31
|
+
so epoch-to-epoch variation is preserved while runs remain
|
|
32
|
+
repeatable).
|
|
33
|
+
"""
|
|
34
|
+
|
|
35
|
+
def __init__(self, labels, batch_size, samples_per_class, seed=None):
|
|
36
|
+
labels = np.asarray(labels)
|
|
37
|
+
if batch_size % samples_per_class != 0:
|
|
38
|
+
fatal(
|
|
39
|
+
f"batch_size ({batch_size}) must be divisible by "
|
|
40
|
+
f"samples_per_class ({samples_per_class})", ValueError)
|
|
41
|
+
|
|
42
|
+
self._batch_size = batch_size
|
|
43
|
+
self._samples_per_class = samples_per_class
|
|
44
|
+
self._n_groups = batch_size // samples_per_class
|
|
45
|
+
|
|
46
|
+
# Build per-class index lists.
|
|
47
|
+
self._class_indices = {}
|
|
48
|
+
for cls in np.unique(labels):
|
|
49
|
+
self._class_indices[int(cls)] = np.where(labels == cls)[0]
|
|
50
|
+
|
|
51
|
+
n_classes = len(self._class_indices)
|
|
52
|
+
if n_classes < self._n_groups:
|
|
53
|
+
logging.warning(f"BalancedBatchSampler: only {n_classes} classes available "
|
|
54
|
+
f"but batch needs {self._n_groups} groups. Some classes "
|
|
55
|
+
f"will be oversampled within each batch.")
|
|
56
|
+
|
|
57
|
+
self._all_classes = np.array(list(self._class_indices.keys()))
|
|
58
|
+
self._n_batches = len(labels) // batch_size
|
|
59
|
+
self._seed = seed
|
|
60
|
+
logging.info(f"BalancedBatchSampler: {len(labels):,} samples, "
|
|
61
|
+
f"{n_classes} classes, batch_size={batch_size}, "
|
|
62
|
+
f"samples_per_class={samples_per_class}, "
|
|
63
|
+
f"batches/epoch={self._n_batches}")
|
|
64
|
+
|
|
65
|
+
def __len__(self):
|
|
66
|
+
return self._n_batches
|
|
67
|
+
|
|
68
|
+
def __iter__(self):
|
|
69
|
+
rng = np.random.default_rng(self._seed)
|
|
70
|
+
for _ in range(self._n_batches):
|
|
71
|
+
# Pick which classes appear in this batch.
|
|
72
|
+
chosen = rng.choice(self._all_classes, size=self._n_groups, replace=True)
|
|
73
|
+
batch = []
|
|
74
|
+
for cls in chosen:
|
|
75
|
+
pool = self._class_indices[int(cls)]
|
|
76
|
+
chosen_idx = rng.choice(
|
|
77
|
+
pool,
|
|
78
|
+
size=self._samples_per_class,
|
|
79
|
+
replace=len(pool) < self._samples_per_class,
|
|
80
|
+
)
|
|
81
|
+
batch.extend(chosen_idx.tolist())
|
|
82
|
+
yield batch
|
|
@@ -0,0 +1,355 @@
|
|
|
1
|
+
"""Unified image dataset for self-supervised pre-training.
|
|
2
|
+
|
|
3
|
+
Stitches together multiple HuggingFace datasets (or pre-downloaded image
|
|
4
|
+
directories) into a single :class:`torch.utils.data.Dataset` with flat
|
|
5
|
+
indexing. All ``__getitem__`` calls return ``dict[str, Any]`` with named
|
|
6
|
+
fields (e.g. ``"image"``, ``"label"``).
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
import bisect
|
|
10
|
+
import logging
|
|
11
|
+
|
|
12
|
+
import numpy as np
|
|
13
|
+
from PIL import Image
|
|
14
|
+
|
|
15
|
+
from genml_kit.datasets.hf_proxy import HFDatasetProxy
|
|
16
|
+
from genml_kit.datasets.image_folder import ImageFolderDataset
|
|
17
|
+
from genml_kit.datasets.retry import getitem_retry
|
|
18
|
+
from genml_kit.utils.logging import fatal
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class _HFDataset:
|
|
22
|
+
"""Wrapper around a HuggingFace ``datasets.Dataset``."""
|
|
23
|
+
|
|
24
|
+
def __init__(self,
|
|
25
|
+
name,
|
|
26
|
+
split="train",
|
|
27
|
+
image_column=None,
|
|
28
|
+
label_column=None,
|
|
29
|
+
cache_dir=None,
|
|
30
|
+
hf_token=None):
|
|
31
|
+
self.name = name
|
|
32
|
+
self._image_column = image_column
|
|
33
|
+
self._label_col = None
|
|
34
|
+
self._label_names = []
|
|
35
|
+
self._label2id = {}
|
|
36
|
+
self._column_map = {}
|
|
37
|
+
from datasets import load_dataset
|
|
38
|
+
logging.info(f"Loading HF dataset '{name}' (split={split}) ...")
|
|
39
|
+
ds = load_dataset(name, split=split, cache_dir=cache_dir, token=hf_token)
|
|
40
|
+
if len(ds) == 0:
|
|
41
|
+
fatal(f"Dataset '{name}' (split={split}) has 0 samples", ValueError)
|
|
42
|
+
col = self._image_column or HFDatasetProxy.detect_image_column(ds)
|
|
43
|
+
if col is None:
|
|
44
|
+
fatal(
|
|
45
|
+
f"No image column detected in '{name}'. "
|
|
46
|
+
f"Columns: {ds.column_names}. Set 'image_column' explicitly.",
|
|
47
|
+
ValueError,
|
|
48
|
+
)
|
|
49
|
+
self._image_column = col
|
|
50
|
+
self._ds = HFDatasetProxy.normalize_image_column(ds, self._image_column)
|
|
51
|
+
self._detect_labels(label_column)
|
|
52
|
+
logging.info(f" '{name}': {len(self._ds):,} samples "
|
|
53
|
+
f"(image column: '{self._image_column}')")
|
|
54
|
+
|
|
55
|
+
def _detect_labels(self, label_column=None):
|
|
56
|
+
"""Detect and normalize the label column, if any."""
|
|
57
|
+
import datasets as ds_lib
|
|
58
|
+
self._label_names = []
|
|
59
|
+
self._label2id = {}
|
|
60
|
+
self._label_col = label_column
|
|
61
|
+
if self._label_col is None:
|
|
62
|
+
self._label_col = HFDatasetProxy.detect_label_column(self._ds)
|
|
63
|
+
self._build_column_map()
|
|
64
|
+
if self._label_col is None:
|
|
65
|
+
return
|
|
66
|
+
if self._label_col not in self._ds.column_names:
|
|
67
|
+
fatal(
|
|
68
|
+
f"Label column '{self._label_col}' is not present. "
|
|
69
|
+
f"Available columns: {self._ds.column_names}",
|
|
70
|
+
ValueError,
|
|
71
|
+
)
|
|
72
|
+
feat = self._ds.features[self._label_col]
|
|
73
|
+
if isinstance(feat, ds_lib.ClassLabel):
|
|
74
|
+
self._label_names = feat.names
|
|
75
|
+
self._label2id = {name: i for i, name in enumerate(self._label_names)}
|
|
76
|
+
else:
|
|
77
|
+
self._label_names = sorted(set(self._ds[self._label_col]))
|
|
78
|
+
self._label2id = {name: i for i, name in enumerate(self._label_names)}
|
|
79
|
+
logging.info(f" '{self.name}': {len(self._label_names)} classes "
|
|
80
|
+
f"(label column: '{self._label_col}')")
|
|
81
|
+
|
|
82
|
+
@property
|
|
83
|
+
def has_labels(self):
|
|
84
|
+
return self._label_col is not None
|
|
85
|
+
|
|
86
|
+
@property
|
|
87
|
+
def num_labels(self):
|
|
88
|
+
return len(self._label_names)
|
|
89
|
+
|
|
90
|
+
@property
|
|
91
|
+
def label_names(self):
|
|
92
|
+
return list(self._label_names)
|
|
93
|
+
|
|
94
|
+
@property
|
|
95
|
+
def column_map(self):
|
|
96
|
+
"""``{SOURCE_COL: ENSEMBLE_COL}`` mapping applied by ``__getitem__``."""
|
|
97
|
+
return dict(self._column_map)
|
|
98
|
+
|
|
99
|
+
def _build_column_map(self):
|
|
100
|
+
"""Build the source-to-ensemble output column mapping.
|
|
101
|
+
|
|
102
|
+
The ensemble unilaterally names its outputs ``"image"`` and
|
|
103
|
+
``"label"``; sub-datasets remap their own column names onto
|
|
104
|
+
those keys.
|
|
105
|
+
"""
|
|
106
|
+
column_map = {self._image_column: "image"}
|
|
107
|
+
if self._label_col is not None:
|
|
108
|
+
column_map[self._label_col] = "label"
|
|
109
|
+
self._column_map = column_map
|
|
110
|
+
|
|
111
|
+
@property
|
|
112
|
+
def labels_array(self):
|
|
113
|
+
"""Numpy array of dataset-local integer labels for all samples."""
|
|
114
|
+
if self._label_col is None:
|
|
115
|
+
return np.array([], dtype=np.int64)
|
|
116
|
+
raw = np.asarray(self._ds[self._label_col])
|
|
117
|
+
mapped = []
|
|
118
|
+
for value in raw:
|
|
119
|
+
if isinstance(value, str):
|
|
120
|
+
if value not in self._label2id:
|
|
121
|
+
fatal(
|
|
122
|
+
f"Dataset '{self.name}': unknown label value {value!r} in "
|
|
123
|
+
f"column '{self._label_col}'. Known labels: "
|
|
124
|
+
f"{sorted(self._label2id)}", ValueError)
|
|
125
|
+
mapped.append(self._label2id[value])
|
|
126
|
+
else:
|
|
127
|
+
label_id = int(value)
|
|
128
|
+
if label_id < 0 or label_id >= len(self._label_names):
|
|
129
|
+
fatal(
|
|
130
|
+
f"Dataset '{self.name}': local label id {label_id} is out "
|
|
131
|
+
f"of range for {len(self._label_names)} classes", ValueError)
|
|
132
|
+
mapped.append(label_id)
|
|
133
|
+
return np.array(mapped, dtype=np.int64)
|
|
134
|
+
|
|
135
|
+
def __len__(self):
|
|
136
|
+
return len(self._ds)
|
|
137
|
+
|
|
138
|
+
def __getitem__(self, idx):
|
|
139
|
+
|
|
140
|
+
def load(i):
|
|
141
|
+
row = self._ds[i]
|
|
142
|
+
image = row[self._image_column]
|
|
143
|
+
if not isinstance(image, Image.Image):
|
|
144
|
+
image = Image.open(image)
|
|
145
|
+
image = image.convert("RGB")
|
|
146
|
+
result = {self._image_column: image}
|
|
147
|
+
if self._label_col is not None:
|
|
148
|
+
raw_label = row[self._label_col]
|
|
149
|
+
result[self._label_col] = self._label2id.get(raw_label, raw_label)
|
|
150
|
+
return result
|
|
151
|
+
|
|
152
|
+
row, _ = getitem_retry(idx, load, len(self._ds))
|
|
153
|
+
return {self._column_map[k]: v for k, v in row.items()}
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
class DatasetEnsemble:
|
|
157
|
+
"""Concatenation of multiple image datasets.
|
|
158
|
+
|
|
159
|
+
Images are returned as dicts with at least an ``"image"`` key
|
|
160
|
+
(a :class:`PIL.Image.Image`). When any dataset provides labels,
|
|
161
|
+
a ``"label"`` key (integer) is also present.
|
|
162
|
+
"""
|
|
163
|
+
|
|
164
|
+
def __init__(self, dataset_configs, cache_dir=None, hf_token=None, strict=False):
|
|
165
|
+
self._datasets = []
|
|
166
|
+
self._offsets = None
|
|
167
|
+
self._global_label_names = []
|
|
168
|
+
self._global_label2id = {}
|
|
169
|
+
|
|
170
|
+
for cfg in dataset_configs:
|
|
171
|
+
name = cfg["name"]
|
|
172
|
+
source = cfg.get("source", "hf")
|
|
173
|
+
try:
|
|
174
|
+
if source == "hf":
|
|
175
|
+
ds = _HFDataset(
|
|
176
|
+
name=name,
|
|
177
|
+
split=cfg.get("split", "train"),
|
|
178
|
+
image_column=cfg.get("image_column"),
|
|
179
|
+
label_column=cfg.get("label_column"),
|
|
180
|
+
cache_dir=cache_dir,
|
|
181
|
+
hf_token=hf_token,
|
|
182
|
+
)
|
|
183
|
+
elif source == "imagefolder":
|
|
184
|
+
ds = ImageFolderDataset(root_dir=name)
|
|
185
|
+
else:
|
|
186
|
+
logging.warning(f"Unknown source '{source}' for dataset '{name}', skipping")
|
|
187
|
+
continue
|
|
188
|
+
n = len(ds)
|
|
189
|
+
if n == 0:
|
|
190
|
+
logging.warning(f"Dataset '{name}' has 0 images, skipping")
|
|
191
|
+
continue
|
|
192
|
+
logging.info(f" + {name}: {n:,} images")
|
|
193
|
+
self._datasets.append(ds)
|
|
194
|
+
except Exception:
|
|
195
|
+
if strict:
|
|
196
|
+
fatal(f"Failed to initialize dataset '{name}' ({source})", RuntimeError)
|
|
197
|
+
logging.exception(
|
|
198
|
+
" Failed to initialize dataset '%s' (%s); skipping.",
|
|
199
|
+
name,
|
|
200
|
+
source,
|
|
201
|
+
)
|
|
202
|
+
|
|
203
|
+
if not self._datasets:
|
|
204
|
+
fatal("No datasets loaded successfully", RuntimeError)
|
|
205
|
+
|
|
206
|
+
self._build_offsets()
|
|
207
|
+
|
|
208
|
+
if self.has_labels:
|
|
209
|
+
self._build_global_label_space()
|
|
210
|
+
|
|
211
|
+
def _build_offsets(self):
|
|
212
|
+
"""Compute cumulative offsets for flat indexing."""
|
|
213
|
+
offsets = [0]
|
|
214
|
+
for ds in self._datasets:
|
|
215
|
+
offsets.append(offsets[-1] + len(ds))
|
|
216
|
+
self._offsets = offsets
|
|
217
|
+
|
|
218
|
+
@property
|
|
219
|
+
def image_column(self):
|
|
220
|
+
"""Canonical image key emitted by every sub-dataset."""
|
|
221
|
+
return "image"
|
|
222
|
+
|
|
223
|
+
@property
|
|
224
|
+
def label_column(self):
|
|
225
|
+
"""Canonical label key, or ``None`` when nothing is labeled."""
|
|
226
|
+
return "label" if self.has_labels else None
|
|
227
|
+
|
|
228
|
+
@property
|
|
229
|
+
def unlabeled_datasets(self):
|
|
230
|
+
"""Names of sub-datasets that provide no labels."""
|
|
231
|
+
return [ds.name for ds in self._datasets if not ds.has_labels]
|
|
232
|
+
|
|
233
|
+
def _build_global_label_space(self):
|
|
234
|
+
"""Map per-dataset label names to a shared integer space.
|
|
235
|
+
|
|
236
|
+
Idempotent: rebuilding is a no-op once the global space exists.
|
|
237
|
+
"""
|
|
238
|
+
if not self._global_label_names:
|
|
239
|
+
all_names = set()
|
|
240
|
+
for ds in self._datasets:
|
|
241
|
+
if ds.has_labels:
|
|
242
|
+
all_names.update(ds.label_names)
|
|
243
|
+
self._global_label_names = sorted(all_names)
|
|
244
|
+
self._global_label2id = {
|
|
245
|
+
name: i for i, name in enumerate(self._global_label_names)
|
|
246
|
+
}
|
|
247
|
+
|
|
248
|
+
def ensure_label_space(self):
|
|
249
|
+
"""Build the global label space across labeled sub-datasets.
|
|
250
|
+
|
|
251
|
+
Mixed ensembles (some sources labeled, some not) are supported:
|
|
252
|
+
unlabeled datasets are skipped with a warning. Public entry
|
|
253
|
+
point for callers that iterate a label-requiring dataset outside
|
|
254
|
+
:meth:`__init__`. Idempotent, so it is cheap to call
|
|
255
|
+
defensively.
|
|
256
|
+
"""
|
|
257
|
+
unlabeled = self.unlabeled_datasets
|
|
258
|
+
if unlabeled:
|
|
259
|
+
logging.warning(
|
|
260
|
+
"ensure_label_space: %d of %d datasets provide no labels and "
|
|
261
|
+
"will be skipped for label-based sampling: %s", len(unlabeled),
|
|
262
|
+
len(self._datasets), ", ".join(unlabeled))
|
|
263
|
+
self._build_global_label_space()
|
|
264
|
+
|
|
265
|
+
def _remap_label(self, ds, local_label):
|
|
266
|
+
"""Convert a dataset-local label id to the global integer id."""
|
|
267
|
+
if isinstance(local_label, bool) or not isinstance(local_label, (int, np.integer)):
|
|
268
|
+
fatal(
|
|
269
|
+
f"Dataset '{ds.name}': expected an integer label id, got "
|
|
270
|
+
f"{local_label!r} ({type(local_label).__name__})", ValueError)
|
|
271
|
+
label_id = int(local_label)
|
|
272
|
+
if label_id < 0 or label_id >= len(ds.label_names):
|
|
273
|
+
fatal(
|
|
274
|
+
f"Dataset '{ds.name}': local label id {label_id} is out of "
|
|
275
|
+
f"range for {len(ds.label_names)} classes", ValueError)
|
|
276
|
+
return self._global_label2id[ds.label_names[label_id]]
|
|
277
|
+
|
|
278
|
+
@property
|
|
279
|
+
def has_labels(self):
|
|
280
|
+
return any(ds.has_labels for ds in self._datasets)
|
|
281
|
+
|
|
282
|
+
@property
|
|
283
|
+
def num_labels(self):
|
|
284
|
+
if not self.has_labels:
|
|
285
|
+
return 0
|
|
286
|
+
return len(self._global_label_names)
|
|
287
|
+
|
|
288
|
+
@property
|
|
289
|
+
def label_names(self):
|
|
290
|
+
if not self.has_labels:
|
|
291
|
+
return []
|
|
292
|
+
return list(self._global_label_names)
|
|
293
|
+
|
|
294
|
+
@property
|
|
295
|
+
def label2id(self):
|
|
296
|
+
"""Mapping of global label name to global integer id."""
|
|
297
|
+
return dict(self._global_label2id)
|
|
298
|
+
|
|
299
|
+
@property
|
|
300
|
+
def labels_array(self):
|
|
301
|
+
"""Global integer labels aligned with ``__getitem__`` positions.
|
|
302
|
+
|
|
303
|
+
Requires every sub-dataset to provide labels; fatals otherwise
|
|
304
|
+
(use :attr:`unlabeled_datasets` to inspect mixed ensembles).
|
|
305
|
+
"""
|
|
306
|
+
unlabeled = self.unlabeled_datasets
|
|
307
|
+
if unlabeled:
|
|
308
|
+
fatal(
|
|
309
|
+
"DatasetEnsemble.labels_array requires every dataset to provide "
|
|
310
|
+
f"labels; these do not: {', '.join(unlabeled)}", ValueError)
|
|
311
|
+
parts = []
|
|
312
|
+
for ds in self._datasets:
|
|
313
|
+
local = ds.labels_array
|
|
314
|
+
global_ids = np.array(
|
|
315
|
+
[self._global_label2id[ds.label_names[int(lbl)]] for lbl in local],
|
|
316
|
+
dtype=np.int64,
|
|
317
|
+
)
|
|
318
|
+
parts.append(global_ids)
|
|
319
|
+
return np.concatenate(parts) if parts else np.array([], dtype=np.int64)
|
|
320
|
+
|
|
321
|
+
def __len__(self):
|
|
322
|
+
return self._offsets[-1] if self._offsets else 0
|
|
323
|
+
|
|
324
|
+
def _get_item(self, idx):
|
|
325
|
+
"""Resolve a valid global index and load its image (and label)."""
|
|
326
|
+
ds_idx = bisect.bisect_right(self._offsets, idx) - 1
|
|
327
|
+
local_idx = idx - self._offsets[ds_idx]
|
|
328
|
+
ds = self._datasets[ds_idx]
|
|
329
|
+
item = ds[local_idx]
|
|
330
|
+
if self.has_labels and "label" in item:
|
|
331
|
+
item["label"] = self._remap_label(ds, item["label"])
|
|
332
|
+
return item
|
|
333
|
+
|
|
334
|
+
def __getitem__(self, idx):
|
|
335
|
+
if idx < 0 or idx >= len(self):
|
|
336
|
+
fatal(
|
|
337
|
+
f"Index {idx} out of range for ensemble of length {len(self)}",
|
|
338
|
+
IndexError,
|
|
339
|
+
)
|
|
340
|
+
return self._get_item(idx)
|
|
341
|
+
|
|
342
|
+
def summary(self):
|
|
343
|
+
"""Return a human-readable summary of the ensemble."""
|
|
344
|
+
lines = [(f"DatasetEnsemble: {len(self):,} images from "
|
|
345
|
+
f"{len(self._datasets)} dataset(s)")]
|
|
346
|
+
for i, ds in enumerate(self._datasets):
|
|
347
|
+
tag = f"{len(ds):,} images"
|
|
348
|
+
if ds.has_labels:
|
|
349
|
+
tag += f", {ds.num_labels} classes"
|
|
350
|
+
else:
|
|
351
|
+
tag += ", NO LABELS"
|
|
352
|
+
lines.append(f" [{i}] {ds.name}: {tag}")
|
|
353
|
+
if self.has_labels:
|
|
354
|
+
lines.append(f" Global label space: {self.num_labels} classes")
|
|
355
|
+
return "\n".join(lines)
|
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
"""FieldSectorDataset — select and rename fields from a dict-returning dataset."""
|
|
2
|
+
|
|
3
|
+
from genml_kit.utils.logging import fatal
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class FieldSectorDataset:
|
|
7
|
+
"""Wrap a dict-returning dataset and keep only the selected fields.
|
|
8
|
+
|
|
9
|
+
Parameters
|
|
10
|
+
----------
|
|
11
|
+
source : Dataset
|
|
12
|
+
A dataset whose ``__getitem__`` returns a ``dict``.
|
|
13
|
+
fields : dict[str, str]
|
|
14
|
+
Mapping ``{SRC_NAME: DST_NAME}``. For every item, the value at
|
|
15
|
+
key *SRC_NAME* is picked from the source item and stored under
|
|
16
|
+
*DST_NAME* in the returned dict. All other fields are dropped.
|
|
17
|
+
|
|
18
|
+
Example
|
|
19
|
+
-------
|
|
20
|
+
>>> ds = FieldSectorDataset(my_dataset, fields={"img": "image"})
|
|
21
|
+
>>> ds[0] # {"image": <value of my_dataset[0]["img"]>}
|
|
22
|
+
"""
|
|
23
|
+
|
|
24
|
+
def __init__(self, source, fields):
|
|
25
|
+
if not fields:
|
|
26
|
+
fatal("FieldSectorDataset requires a non-empty 'fields' mapping.", ValueError)
|
|
27
|
+
self._source = source
|
|
28
|
+
self._fields = dict(fields)
|
|
29
|
+
|
|
30
|
+
def __len__(self):
|
|
31
|
+
return len(self._source)
|
|
32
|
+
|
|
33
|
+
@property
|
|
34
|
+
def fields(self):
|
|
35
|
+
"""The ``{SRC_NAME: DST_NAME}`` selection mapping."""
|
|
36
|
+
return dict(self._fields)
|
|
37
|
+
|
|
38
|
+
def __getitem__(self, idx):
|
|
39
|
+
item = self._source[idx]
|
|
40
|
+
return {dst: item[src] for src, dst in self._fields.items()}
|