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.
Files changed (94) hide show
  1. genml_kit/__init__.py +11 -0
  2. genml_kit/datasets/__init__.py +11 -0
  3. genml_kit/datasets/balanced_sampler.py +82 -0
  4. genml_kit/datasets/ensemble.py +355 -0
  5. genml_kit/datasets/field_dataset.py +40 -0
  6. genml_kit/datasets/hf_proxy.py +239 -0
  7. genml_kit/datasets/image_folder.py +264 -0
  8. genml_kit/datasets/retry.py +43 -0
  9. genml_kit/datasets/weighted_sampler.py +70 -0
  10. genml_kit/io/__init__.py +1 -0
  11. genml_kit/io/checkpointing.py +721 -0
  12. genml_kit/io/storage_utils.py +339 -0
  13. genml_kit/models/__init__.py +37 -0
  14. genml_kit/models/attention_pooling.py +43 -0
  15. genml_kit/models/byol.py +98 -0
  16. genml_kit/models/cls_model_wrapper/__init__.py +9 -0
  17. genml_kit/models/cls_model_wrapper/loader.py +50 -0
  18. genml_kit/models/cls_model_wrapper/model.py +102 -0
  19. genml_kit/models/contrastive.py +74 -0
  20. genml_kit/models/convvit/__init__.py +13 -0
  21. genml_kit/models/convvit/loader.py +162 -0
  22. genml_kit/models/convvit/masked_encoder.py +32 -0
  23. genml_kit/models/convvit/model.py +209 -0
  24. genml_kit/models/convvit/processor.py +47 -0
  25. genml_kit/models/dino.py +151 -0
  26. genml_kit/models/encoder_utils.py +83 -0
  27. genml_kit/models/masked_encoder.py +45 -0
  28. genml_kit/models/processors/__init__.py +5 -0
  29. genml_kit/models/processors/base.py +76 -0
  30. genml_kit/models/registry.py +231 -0
  31. genml_kit/models/simmim.py +122 -0
  32. genml_kit/models/timm/__init__.py +3 -0
  33. genml_kit/models/timm/loader.py +110 -0
  34. genml_kit/models/timm/model.py +59 -0
  35. genml_kit/models/timm/processor.py +61 -0
  36. genml_kit/models/transformer_utils.py +86 -0
  37. genml_kit/models/uvito/__init__.py +11 -0
  38. genml_kit/models/uvito/loader.py +116 -0
  39. genml_kit/models/uvito/model.py +150 -0
  40. genml_kit/models/uvito/processor.py +33 -0
  41. genml_kit/pretrain/__init__.py +1 -0
  42. genml_kit/pretrain/augmentations/__init__.py +1 -0
  43. genml_kit/pretrain/augmentations/dual_view.py +11 -0
  44. genml_kit/pretrain/augmentations/multicrop.py +105 -0
  45. genml_kit/pretrain/cli.py +884 -0
  46. genml_kit/pretrain/losses/__init__.py +8 -0
  47. genml_kit/pretrain/losses/byol.py +24 -0
  48. genml_kit/pretrain/losses/contrastive.py +54 -0
  49. genml_kit/pretrain/losses/dino.py +82 -0
  50. genml_kit/pretrain/losses/focal.py +77 -0
  51. genml_kit/pretrain/methods/__init__.py +11 -0
  52. genml_kit/pretrain/methods/base.py +100 -0
  53. genml_kit/pretrain/methods/byol.py +105 -0
  54. genml_kit/pretrain/methods/dino.py +179 -0
  55. genml_kit/pretrain/methods/ijepa.py +310 -0
  56. genml_kit/pretrain/methods/registry.py +33 -0
  57. genml_kit/pretrain/methods/simmim.py +138 -0
  58. genml_kit/pretrain/methods/supcon.py +72 -0
  59. genml_kit/training/__init__.py +1 -0
  60. genml_kit/training/classifiers/__init__.py +80 -0
  61. genml_kit/training/classifiers/base.py +45 -0
  62. genml_kit/training/classifiers/cls_attention.py +91 -0
  63. genml_kit/training/classifiers/mlp.py +54 -0
  64. genml_kit/training/eval.py +128 -0
  65. genml_kit/training/grad_monitor.py +434 -0
  66. genml_kit/training/infer.py +222 -0
  67. genml_kit/training/metrics.py +76 -0
  68. genml_kit/training/model_utils.py +456 -0
  69. genml_kit/training/optim_factory.py +459 -0
  70. genml_kit/training/param_align.py +166 -0
  71. genml_kit/training/train.py +1381 -0
  72. genml_kit/training/train_reporting.py +197 -0
  73. genml_kit/training/tta.py +91 -0
  74. genml_kit/training/xgb_pipeline.py +153 -0
  75. genml_kit/training/xgb_utils.py +144 -0
  76. genml_kit/utils/__init__.py +1 -0
  77. genml_kit/utils/args.py +176 -0
  78. genml_kit/utils/attr.py +141 -0
  79. genml_kit/utils/cli.py +108 -0
  80. genml_kit/utils/gpu.py +41 -0
  81. genml_kit/utils/image_dump.py +110 -0
  82. genml_kit/utils/label.py +27 -0
  83. genml_kit/utils/logging.py +150 -0
  84. genml_kit/utils/script.py +82 -0
  85. genml_kit/utils/seed.py +86 -0
  86. genml_kit/utils/signal.py +114 -0
  87. genml_kit/utils/table.py +58 -0
  88. genml_kit/utils/transformer.py +57 -0
  89. genml_kit-0.1.0.dist-info/METADATA +1209 -0
  90. genml_kit-0.1.0.dist-info/RECORD +94 -0
  91. genml_kit-0.1.0.dist-info/WHEEL +5 -0
  92. genml_kit-0.1.0.dist-info/entry_points.txt +4 -0
  93. genml_kit-0.1.0.dist-info/licenses/LICENSE +190 -0
  94. 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()}