netloader 3.11.2__tar.gz → 3.12.0__tar.gz
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.
- {netloader-3.11.2 → netloader-3.12.0}/PKG-INFO +1 -1
- {netloader-3.11.2 → netloader-3.12.0}/netloader/__init__.py +1 -1
- {netloader-3.11.2 → netloader-3.12.0}/netloader/architectures/base.py +13 -1
- {netloader-3.11.2 → netloader-3.12.0}/netloader/architectures/encoder_decoder.py +15 -1
- {netloader-3.11.2 → netloader-3.12.0}/netloader/architectures/flows.py +13 -5
- netloader-3.12.0/netloader/data/__init__.py +8 -0
- netloader-3.12.0/netloader/data/datasets.py +270 -0
- netloader-3.11.2/netloader/data.py → netloader-3.12.0/netloader/data/structures.py +219 -310
- {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/base.py +5 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/misc.py +5 -5
- {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/multi_layer.py +1 -1
- {netloader-3.11.2 → netloader-3.12.0}/netloader/models/misc.py +27 -15
- {netloader-3.11.2 → netloader-3.12.0}/netloader/network.py +91 -14
- {netloader-3.11.2 → netloader-3.12.0}/netloader/transforms.py +23 -5
- {netloader-3.11.2 → netloader-3.12.0}/netloader/utils/utils.py +36 -7
- {netloader-3.11.2 → netloader-3.12.0}/netloader.egg-info/PKG-INFO +1 -1
- {netloader-3.11.2 → netloader-3.12.0}/netloader.egg-info/SOURCES.txt +3 -1
- {netloader-3.11.2 → netloader-3.12.0}/LICENSE.txt +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/README.md +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/architectures/__init__.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/architectures/utils.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/__init__.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/blocks.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/convolutional.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/flows.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/linear.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/pooling.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/recurrent.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/utils.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/loss_funcs.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/models/__init__.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/models/convnext.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/networks/__init__.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/schedulers.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/utils/__init__.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/utils/configs.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/utils/transforms.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader/utils/types.py +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader.egg-info/dependency_links.txt +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader.egg-info/requires.txt +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/netloader.egg-info/top_level.txt +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/pyproject.toml +0 -0
- {netloader-3.11.2 → netloader-3.12.0}/setup.cfg +0 -0
|
@@ -6,7 +6,7 @@ import logging
|
|
|
6
6
|
import warnings
|
|
7
7
|
|
|
8
8
|
|
|
9
|
-
__version__ = '3.
|
|
9
|
+
__version__ = '3.12.0'
|
|
10
10
|
__author__ = 'Ethan Tregidga'
|
|
11
11
|
logging.basicConfig(format='%(levelname)s: %(message)s', level=logging.WARNING)
|
|
12
12
|
warnings.filterwarnings('once', category=DeprecationWarning, module=r'^netloader(\.|$)')
|
|
@@ -46,6 +46,8 @@ class BaseArchitecture(UtilityMixin, ABC, Generic[LossCT, TensorLossCT]):
|
|
|
46
46
|
|
|
47
47
|
Attributes
|
|
48
48
|
----------
|
|
49
|
+
checkpoints : bool
|
|
50
|
+
If checkpoints should be saved after network forward pass
|
|
49
51
|
description : str
|
|
50
52
|
Description of the architecture
|
|
51
53
|
version : str
|
|
@@ -62,6 +64,8 @@ class BaseArchitecture(UtilityMixin, ABC, Generic[LossCT, TensorLossCT]):
|
|
|
62
64
|
Architecture optimiser
|
|
63
65
|
scheduler : LRScheduler
|
|
64
66
|
Optimiser scheduler
|
|
67
|
+
run : Run | None
|
|
68
|
+
Weights & Biases run object for logging and tracking experiments
|
|
65
69
|
net : BaseNetwork
|
|
66
70
|
Neural network
|
|
67
71
|
"""
|
|
@@ -79,6 +83,7 @@ class BaseArchitecture(UtilityMixin, ABC, Generic[LossCT, TensorLossCT]):
|
|
|
79
83
|
_scheduler_kwargs: dict[str, Any]
|
|
80
84
|
_logger: log.Logger = log.getLogger(__name__)
|
|
81
85
|
_device: torch.device = torch.device('cpu')
|
|
86
|
+
checkpoints: bool = False
|
|
82
87
|
description: str
|
|
83
88
|
version: str = netloader.__version__
|
|
84
89
|
losses: tuple[list[LossCT], list[LossCT]]
|
|
@@ -283,7 +288,7 @@ class BaseArchitecture(UtilityMixin, ABC, Generic[LossCT, TensorLossCT]):
|
|
|
283
288
|
self.net = state['net']
|
|
284
289
|
self.run = state.get('run', None)
|
|
285
290
|
|
|
286
|
-
if not utils.compare_versions(self.version, netloader.__version__):
|
|
291
|
+
if not utils.compare_versions(self.version, netloader.__version__, level='minor'):
|
|
287
292
|
warn(
|
|
288
293
|
f'Architecture version ({self.version}) is older than the current '
|
|
289
294
|
f'NetLoader version ({netloader.__version__}), please resave the architecture '
|
|
@@ -513,6 +518,9 @@ class BaseArchitecture(UtilityMixin, ABC, Generic[LossCT, TensorLossCT]):
|
|
|
513
518
|
|
|
514
519
|
self._update(loss['total'] if isinstance(loss, dict) else loss)
|
|
515
520
|
|
|
521
|
+
if not self.checkpoints:
|
|
522
|
+
self.net.clear_checkpoints()
|
|
523
|
+
|
|
516
524
|
if isinstance(loss, dict):
|
|
517
525
|
return {key: value.item() for key, value in loss.items()} # type: ignore[return-value]
|
|
518
526
|
return loss.item() # type: ignore[return-value]
|
|
@@ -1067,6 +1075,9 @@ class BaseArchitecture(UtilityMixin, ABC, Generic[LossCT, TensorLossCT]):
|
|
|
1067
1075
|
]))
|
|
1068
1076
|
self._predict_print(i, time() - t_initial, loader)
|
|
1069
1077
|
|
|
1078
|
+
if not self.checkpoints:
|
|
1079
|
+
self.net.clear_checkpoints()
|
|
1080
|
+
|
|
1070
1081
|
# Transforms all data and saves it to a dictionary
|
|
1071
1082
|
for (key, transform), datum_ in zip(transforms.items(), zip(*data)):
|
|
1072
1083
|
if datum_[0] is None:
|
|
@@ -1085,6 +1096,7 @@ class BaseArchitecture(UtilityMixin, ABC, Generic[LossCT, TensorLossCT]):
|
|
|
1085
1096
|
datum,
|
|
1086
1097
|
transform if isinstance(transform, list) else repeat(transform),
|
|
1087
1098
|
)],
|
|
1099
|
+
names=datum.get_names(),
|
|
1088
1100
|
)
|
|
1089
1101
|
elif isinstance(transform, BaseTransform):
|
|
1090
1102
|
assert not isinstance(datum, DataList)
|
|
@@ -27,7 +27,7 @@ class BaseEncoder(BaseArchitecture):
|
|
|
27
27
|
description : str
|
|
28
28
|
Description of the architecture
|
|
29
29
|
version : str
|
|
30
|
-
Version of the architecture
|
|
30
|
+
Version of the architecture when it was created or re-saved
|
|
31
31
|
losses : tuple[list[LossCT], list[LossCT]]
|
|
32
32
|
Architecture training and validation losses as a float or dictionary of losses for each loss
|
|
33
33
|
function
|
|
@@ -42,6 +42,8 @@ class BaseEncoder(BaseArchitecture):
|
|
|
42
42
|
Architecture optimiser
|
|
43
43
|
scheduler : LRScheduler
|
|
44
44
|
Optimiser scheduler
|
|
45
|
+
run : Run | None
|
|
46
|
+
Weights & Biases run object for logging and tracking experiments
|
|
45
47
|
net : BaseNetwork
|
|
46
48
|
Neural network
|
|
47
49
|
"""
|
|
@@ -211,6 +213,8 @@ class Autoencoder(BaseArchitecture):
|
|
|
211
213
|
----------
|
|
212
214
|
description : str
|
|
213
215
|
Description of the architecture
|
|
216
|
+
version : str
|
|
217
|
+
Version of the architecture when it was created or re-saved
|
|
214
218
|
losses : tuple[list[LossCT], list[LossCT]]
|
|
215
219
|
Architecture training and validation losses as a float or dictionary of losses for each loss
|
|
216
220
|
function
|
|
@@ -223,6 +227,8 @@ class Autoencoder(BaseArchitecture):
|
|
|
223
227
|
Architecture optimiser
|
|
224
228
|
scheduler : LRScheduler
|
|
225
229
|
Optimiser scheduler
|
|
230
|
+
run : Run | None
|
|
231
|
+
Weights & Biases run object for logging and tracking experiments
|
|
226
232
|
net : BaseNetwork
|
|
227
233
|
Neural network
|
|
228
234
|
reconstruct_func : BaseLoss
|
|
@@ -427,6 +433,8 @@ class Decoder(BaseArchitecture):
|
|
|
427
433
|
----------
|
|
428
434
|
description : str
|
|
429
435
|
Description of the architecture
|
|
436
|
+
version : str
|
|
437
|
+
Version of the architecture when it was created or re-saved
|
|
430
438
|
losses : tuple[list[LossCT], list[LossCT]]
|
|
431
439
|
Architecture training and validation losses as a float or dictionary of losses for each loss
|
|
432
440
|
function
|
|
@@ -441,6 +449,8 @@ class Decoder(BaseArchitecture):
|
|
|
441
449
|
Architecture optimiser
|
|
442
450
|
scheduler : LRScheduler
|
|
443
451
|
Optimiser scheduler
|
|
452
|
+
run : Run | None
|
|
453
|
+
Weights & Biases run object for logging and tracking experiments
|
|
444
454
|
net : BaseNetwork
|
|
445
455
|
Neural network
|
|
446
456
|
"""
|
|
@@ -587,6 +597,8 @@ class Encoder(BaseEncoder):
|
|
|
587
597
|
----------
|
|
588
598
|
description : str
|
|
589
599
|
Description of the architecture
|
|
600
|
+
version : str
|
|
601
|
+
Version of the architecture when it was created or re-saved
|
|
590
602
|
losses : tuple[list[LossCT], list[LossCT]]
|
|
591
603
|
Architecture training and validation losses as a float or dictionary of losses for each loss
|
|
592
604
|
function
|
|
@@ -601,6 +613,8 @@ class Encoder(BaseEncoder):
|
|
|
601
613
|
Architecture optimiser
|
|
602
614
|
scheduler : LRScheduler
|
|
603
615
|
Optimiser scheduler
|
|
616
|
+
run : Run | None
|
|
617
|
+
Weights & Biases run object for logging and tracking experiments
|
|
604
618
|
net : BaseNetwork
|
|
605
619
|
Neural network
|
|
606
620
|
"""
|
|
@@ -14,8 +14,8 @@ from zuko.distributions import NormalizingFlow
|
|
|
14
14
|
from netloader.data import Data
|
|
15
15
|
from netloader.utils import label_change
|
|
16
16
|
from netloader.loss_funcs import BaseLoss
|
|
17
|
-
from netloader.network import BaseNetwork
|
|
18
17
|
from netloader.transforms import BaseTransform
|
|
18
|
+
from netloader.network import BaseNetwork, Network
|
|
19
19
|
from netloader.architectures.base import BaseArchitecture
|
|
20
20
|
from netloader.architectures.encoder_decoder import BaseEncoder
|
|
21
21
|
from netloader.utils.types import NDArrayLike, TensorListLike, NDArrayListLike, TensorT
|
|
@@ -29,10 +29,10 @@ class NormFlow(BaseArchitecture):
|
|
|
29
29
|
|
|
30
30
|
Attributes
|
|
31
31
|
----------
|
|
32
|
-
net : BaseNetwork
|
|
33
|
-
Neural spline flow
|
|
34
32
|
description : str
|
|
35
33
|
Description of the architecture
|
|
34
|
+
version : str
|
|
35
|
+
Version of the architecture when it was created or re-saved
|
|
36
36
|
losses : tuple[list[LossCT], list[LossCT]]
|
|
37
37
|
Architecture training and validation losses as a float or dictionary of losses for each loss
|
|
38
38
|
function
|
|
@@ -45,6 +45,10 @@ class NormFlow(BaseArchitecture):
|
|
|
45
45
|
Architecture optimiser
|
|
46
46
|
scheduler : LRScheduler
|
|
47
47
|
Optimiser scheduler
|
|
48
|
+
run : Run | None
|
|
49
|
+
Weights & Biases run object for logging and tracking experiments
|
|
50
|
+
net : BaseNetwork
|
|
51
|
+
Neural spline flow
|
|
48
52
|
"""
|
|
49
53
|
@staticmethod
|
|
50
54
|
def _data_loader_translation(low_dim: TensorT, high_dim: TensorT) -> tuple[TensorT, TensorT]:
|
|
@@ -132,6 +136,8 @@ class NormFlowEncoder(BaseEncoder):
|
|
|
132
136
|
----------
|
|
133
137
|
description : str
|
|
134
138
|
Description of the architecture
|
|
139
|
+
version : str
|
|
140
|
+
Version of the architecture when it was created or re-saved
|
|
135
141
|
losses : tuple[list[LossCT], list[LossCT]]
|
|
136
142
|
Architecture training and validation losses as a float or dictionary of losses for each loss
|
|
137
143
|
function
|
|
@@ -146,6 +152,8 @@ class NormFlowEncoder(BaseEncoder):
|
|
|
146
152
|
Architecture optimiser
|
|
147
153
|
scheduler : LRScheduler
|
|
148
154
|
Optimiser scheduler
|
|
155
|
+
run : Run | None
|
|
156
|
+
Weights & Biases run object for logging and tracking experiments
|
|
149
157
|
net : BaseNetwork
|
|
150
158
|
Neural network
|
|
151
159
|
"""
|
|
@@ -222,7 +230,7 @@ class NormFlowEncoder(BaseEncoder):
|
|
|
222
230
|
net,
|
|
223
231
|
overwrite=overwrite,
|
|
224
232
|
mix_precision=mix_precision,
|
|
225
|
-
learning_rate=0
|
|
233
|
+
learning_rate=0.,
|
|
226
234
|
description=description,
|
|
227
235
|
verbose=verbose,
|
|
228
236
|
classes=classes,
|
|
@@ -243,7 +251,7 @@ class NormFlowEncoder(BaseEncoder):
|
|
|
243
251
|
'max': transform,
|
|
244
252
|
'meds': transform,
|
|
245
253
|
}
|
|
246
|
-
assert isinstance(self.net,
|
|
254
|
+
assert isinstance(self.net, Network)
|
|
247
255
|
|
|
248
256
|
if not self._train_encoder:
|
|
249
257
|
self.net.layers[:-1].requires_grad_(False)
|
|
@@ -0,0 +1,8 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Base dataset classes and data structures for NetLoader.
|
|
3
|
+
"""
|
|
4
|
+
from netloader.data.datasets import BaseDataset, loader_init
|
|
5
|
+
from netloader.data.structures import ApplyFunc, Data, DataList, data_collation
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
__all__ = ['ApplyFunc', 'Data', 'DataList', 'BaseDataset', 'data_collation', 'loader_init']
|
|
@@ -0,0 +1,270 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Base dataset classes for use with BaseArchitecture.
|
|
3
|
+
"""
|
|
4
|
+
from __future__ import annotations
|
|
5
|
+
import logging as log
|
|
6
|
+
from typing import Any, Generic, Literal, Sequence, cast, overload
|
|
7
|
+
|
|
8
|
+
import torch
|
|
9
|
+
import numpy as np
|
|
10
|
+
from torch.utils.data import Dataset, DataLoader, Subset
|
|
11
|
+
from numpy import ndarray
|
|
12
|
+
|
|
13
|
+
from netloader.utils.types import ArrayLike, DatasetT, DataListT
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class BaseDatasetMeta(type):
|
|
17
|
+
"""
|
|
18
|
+
Automatically creates an index for each sample in the dataset after the dataset has been
|
|
19
|
+
initialised.
|
|
20
|
+
"""
|
|
21
|
+
def __call__(cls: type[DatasetT], *args: Any, **kwargs: Any) -> DatasetT:
|
|
22
|
+
"""
|
|
23
|
+
Parameters
|
|
24
|
+
----------
|
|
25
|
+
cls : type[DatasetT]
|
|
26
|
+
Class that inherited BaseDatasetMeta
|
|
27
|
+
*args
|
|
28
|
+
Optional arguments to pass to BaseDataset class
|
|
29
|
+
**kwargs
|
|
30
|
+
Optional keyword arguments to pass to BaseDataset class
|
|
31
|
+
|
|
32
|
+
Returns
|
|
33
|
+
-------
|
|
34
|
+
DatasetT
|
|
35
|
+
Dataset instance from the class that inherited BaseDatasetMeta
|
|
36
|
+
"""
|
|
37
|
+
instance: DatasetT = type.__call__(cls, *args, **kwargs)
|
|
38
|
+
|
|
39
|
+
if hasattr(instance, 'idxs') and instance.idxs.dtype != np.int_:
|
|
40
|
+
raise ValueError(f'idxs attribute already exists and does not have type int '
|
|
41
|
+
f'({instance.idxs.dtype}), idxs attribute must be reserved for sample '
|
|
42
|
+
f'index')
|
|
43
|
+
|
|
44
|
+
if not hasattr(instance, 'high_dim') or instance.high_dim is None:
|
|
45
|
+
raise ValueError(f'{instance.__class__.__name__} has no high_dim attribute which is '
|
|
46
|
+
f'required by BaseDatasetMeta for creating idxs attribute')
|
|
47
|
+
|
|
48
|
+
if len(instance.idxs) == 0:
|
|
49
|
+
instance.idxs = np.arange(len(instance.high_dim))
|
|
50
|
+
elif len(instance.idxs) != len(instance.high_dim):
|
|
51
|
+
log.getLogger(__name__).warning(f'Length of idxs ({len(instance.idxs)}) and length of '
|
|
52
|
+
f'high_dim ({len(instance.high_dim)}) does not match, '
|
|
53
|
+
f'idxs will be sent to a range of high_dim length')
|
|
54
|
+
instance.idxs = np.arange(len(instance.high_dim))
|
|
55
|
+
|
|
56
|
+
for attribute in ('extra', 'low_dim', 'high_dim'):
|
|
57
|
+
if (getattr(instance, attribute) is not None and
|
|
58
|
+
len(getattr(instance, attribute)) != len(instance.idxs)):
|
|
59
|
+
raise ValueError(f'Length of attribute {attribute} '
|
|
60
|
+
f'({len(getattr(instance, attribute))}) and idxs '
|
|
61
|
+
f'({len(instance.idxs)}) does not match')
|
|
62
|
+
return instance
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
class BaseDataset(
|
|
66
|
+
Dataset[tuple[int, DataListT, DataListT, Any]],
|
|
67
|
+
Generic[DataListT],
|
|
68
|
+
metaclass=BaseDatasetMeta):
|
|
69
|
+
"""
|
|
70
|
+
Base dataset class for use with BaseNetwork.
|
|
71
|
+
|
|
72
|
+
Attributes
|
|
73
|
+
----------
|
|
74
|
+
extra : list[Any] | ArrayLike | None
|
|
75
|
+
Additional data for each sample in the dataset of length N with shape (N,...) and type Any
|
|
76
|
+
idxs : ndarray
|
|
77
|
+
Index for each sample in the dataset with shape (N) and type int
|
|
78
|
+
low_dim : DataListT | None
|
|
79
|
+
Low dimensional data for each sample in the dataset with shape (N,...)
|
|
80
|
+
high_dim : DataListT | None
|
|
81
|
+
High dimensional data for each sample in the dataset with shape (N,...), this is required
|
|
82
|
+
"""
|
|
83
|
+
def __init__(self) -> None:
|
|
84
|
+
super().__init__()
|
|
85
|
+
self.extra: list[Any] | ArrayLike | None = None
|
|
86
|
+
self.idxs: ndarray = np.array([], dtype=np.int_)
|
|
87
|
+
self.low_dim: DataListT | None = None
|
|
88
|
+
self.high_dim: DataListT | None = None
|
|
89
|
+
|
|
90
|
+
def __len__(self) -> int:
|
|
91
|
+
"""
|
|
92
|
+
Returns the number of samples in the dataset
|
|
93
|
+
|
|
94
|
+
Returns
|
|
95
|
+
-------
|
|
96
|
+
int
|
|
97
|
+
Number of samples in the dataset
|
|
98
|
+
"""
|
|
99
|
+
return len(self.idxs)
|
|
100
|
+
|
|
101
|
+
def __getitem__(self, idx: int) -> tuple[int, DataListT, DataListT, Any]:
|
|
102
|
+
"""
|
|
103
|
+
Parameters
|
|
104
|
+
----------
|
|
105
|
+
idx : int
|
|
106
|
+
Sample index
|
|
107
|
+
|
|
108
|
+
Returns
|
|
109
|
+
-------
|
|
110
|
+
tuple[int, DataListT, DataListT, Any]
|
|
111
|
+
Sample index, low dimensional data, high dimensional data, and extra data
|
|
112
|
+
"""
|
|
113
|
+
return self.idxs[idx], self.get_low_dim(idx), self.get_high_dim(idx), self.get_extra(idx)
|
|
114
|
+
|
|
115
|
+
def get_extra(self, idx: int) -> Any:
|
|
116
|
+
"""
|
|
117
|
+
Gets extra data for the sample of the given index
|
|
118
|
+
|
|
119
|
+
Parameters
|
|
120
|
+
----------
|
|
121
|
+
idx : int
|
|
122
|
+
Sample index
|
|
123
|
+
|
|
124
|
+
Returns
|
|
125
|
+
-------
|
|
126
|
+
Any
|
|
127
|
+
Sample extra data
|
|
128
|
+
"""
|
|
129
|
+
return torch.tensor(()) if self.extra is None else self.extra[idx]
|
|
130
|
+
|
|
131
|
+
def get_high_dim(self, idx: int) -> DataListT:
|
|
132
|
+
"""
|
|
133
|
+
Gets a high dimensional sample of the given index
|
|
134
|
+
|
|
135
|
+
Parameters
|
|
136
|
+
----------
|
|
137
|
+
idx : int
|
|
138
|
+
Sample index
|
|
139
|
+
|
|
140
|
+
Returns
|
|
141
|
+
-------
|
|
142
|
+
DataListT
|
|
143
|
+
High dimensional sample
|
|
144
|
+
"""
|
|
145
|
+
assert self.high_dim is not None
|
|
146
|
+
return cast(DataListT, self.high_dim[idx])
|
|
147
|
+
|
|
148
|
+
def get_low_dim(self, idx: int) -> DataListT:
|
|
149
|
+
"""
|
|
150
|
+
Gets a low dimensional sample of the given index
|
|
151
|
+
|
|
152
|
+
Parameters
|
|
153
|
+
----------
|
|
154
|
+
idx : int
|
|
155
|
+
Sample index
|
|
156
|
+
|
|
157
|
+
Returns
|
|
158
|
+
-------
|
|
159
|
+
DataListT
|
|
160
|
+
Low dimensional sample
|
|
161
|
+
"""
|
|
162
|
+
return cast(DataListT, torch.tensor(()) if self.low_dim is None else self.low_dim[idx])
|
|
163
|
+
|
|
164
|
+
def step(self, epoch: float) -> None:
|
|
165
|
+
"""
|
|
166
|
+
Step the dataset for each training iteration.
|
|
167
|
+
|
|
168
|
+
Parameters
|
|
169
|
+
----------
|
|
170
|
+
epoch : float
|
|
171
|
+
Current epoch number
|
|
172
|
+
"""
|
|
173
|
+
|
|
174
|
+
|
|
175
|
+
@overload
|
|
176
|
+
def loader_init(
|
|
177
|
+
dataset: DatasetT,
|
|
178
|
+
*,
|
|
179
|
+
return_idxs: Literal[False],
|
|
180
|
+
batch_size: int = ...,
|
|
181
|
+
ratios: list[float] | tuple[float, ...] | None = ...,
|
|
182
|
+
idxs: list[ndarray] | tuple[ndarray, ...] | ndarray | None = ...,
|
|
183
|
+
**kwargs: Any) -> tuple[DataLoader[tuple[int, DataListT, DataListT, Any]], ...]: ...
|
|
184
|
+
|
|
185
|
+
@overload
|
|
186
|
+
def loader_init(
|
|
187
|
+
dataset: DatasetT,
|
|
188
|
+
*,
|
|
189
|
+
return_idxs: Literal[True],
|
|
190
|
+
batch_size: int = ...,
|
|
191
|
+
ratios: list[float] | tuple[float, ...] | None = ...,
|
|
192
|
+
idxs: list[ndarray] | tuple[ndarray, ...] | ndarray | None = ...,
|
|
193
|
+
**kwargs: Any) -> tuple[
|
|
194
|
+
tuple[DataLoader[tuple[int, DataListT, DataListT, Any]], ...],
|
|
195
|
+
tuple[Sequence[int], ...]]: ...
|
|
196
|
+
|
|
197
|
+
def loader_init(
|
|
198
|
+
dataset: DatasetT,
|
|
199
|
+
*,
|
|
200
|
+
return_idxs: bool = False,
|
|
201
|
+
batch_size: int = 64,
|
|
202
|
+
ratios: list[float] | tuple[float, ...] | None = None,
|
|
203
|
+
idxs: list[ndarray] | tuple[ndarray, ...] | ndarray | None = None,
|
|
204
|
+
**kwargs: Any,
|
|
205
|
+
) -> (tuple[DataLoader[tuple[int, DataListT, DataListT, Any]], ...] |
|
|
206
|
+
tuple[
|
|
207
|
+
tuple[DataLoader[tuple[int, DataListT, DataListT, Any]], ...],
|
|
208
|
+
tuple[Sequence[int], ...]]):
|
|
209
|
+
# pylint: disable=line-too-long
|
|
210
|
+
"""
|
|
211
|
+
Initialises data loaders from a subset of the dataset with the given ratios.
|
|
212
|
+
|
|
213
|
+
Parameters
|
|
214
|
+
----------
|
|
215
|
+
dataset : DatasetT
|
|
216
|
+
Dataset to create data loaders from
|
|
217
|
+
return_idxs : bool, Optional
|
|
218
|
+
If the indexes for each data loader should be returned, default = False
|
|
219
|
+
batch_size : int, Optional
|
|
220
|
+
Batch size when sampling from the data loaders, default = 64
|
|
221
|
+
ratios : list[float] | tuple[float, ...] | None, Optional
|
|
222
|
+
Ratios of length M to split up the dataset into subsets, if idxs is provided, dataset will
|
|
223
|
+
first be split up using idxs and ratios will be used on the remaining samples,
|
|
224
|
+
default = (0.8,0.2)
|
|
225
|
+
idxs : list[ndarray] | tuple[ndarray, ...] | ndarray | None, Optional
|
|
226
|
+
Dataset indexes for creating the subsets with shape (N,S), where N is the number of subsets
|
|
227
|
+
and S is the number of samples in each subset
|
|
228
|
+
**kwargs : Any
|
|
229
|
+
Optional keyword arguments to pass to DataLoader
|
|
230
|
+
|
|
231
|
+
Returns
|
|
232
|
+
-------
|
|
233
|
+
tuple[DataLoader[tuple[int, DataListT, DataListT, Any]], ...] | tuple[tuple[DataLoader[tuple[int, DataListT, DataListT, Any]], ...], tuple[Sequence[int], ...]]
|
|
234
|
+
Data loaders for each subset of length N + M and optionally indexes used for each subset
|
|
235
|
+
"""
|
|
236
|
+
# pylint: enable=line-too-long
|
|
237
|
+
num: int
|
|
238
|
+
slice_: float | ndarray
|
|
239
|
+
loaders: list[DataLoader[tuple[int, DataListT, DataListT, Any]]] = []
|
|
240
|
+
data_idxs: ndarray = np.arange(len(dataset))
|
|
241
|
+
loader: DataLoader[tuple[int, DataListT, DataListT, Any]]
|
|
242
|
+
idxs = [] if idxs is None else list(idxs) \
|
|
243
|
+
if isinstance(idxs, (tuple, list)) or np.ndim(idxs) > 1 else [idxs]
|
|
244
|
+
ratios = (0.8, 1) if ratios is None else list(np.cumsum(np.array(ratios) / np.sum(ratios)))
|
|
245
|
+
np.random.shuffle(data_idxs)
|
|
246
|
+
|
|
247
|
+
for slice_ in idxs + list(ratios):
|
|
248
|
+
if isinstance(slice_, (int, float)):
|
|
249
|
+
num = max(int(len(data_idxs) * slice_), 1)
|
|
250
|
+
slice_ = data_idxs[:num]
|
|
251
|
+
|
|
252
|
+
if not np.isin(data_idxs, slice_).any():
|
|
253
|
+
continue
|
|
254
|
+
|
|
255
|
+
loaders.append(DataLoader(
|
|
256
|
+
Subset(
|
|
257
|
+
cast(Dataset[tuple[int, DataListT, DataListT, Any]], dataset),
|
|
258
|
+
data_idxs[np.isin(data_idxs, slice_)].tolist(),
|
|
259
|
+
),
|
|
260
|
+
batch_size=batch_size,
|
|
261
|
+
**{'shuffle': True} | kwargs,
|
|
262
|
+
))
|
|
263
|
+
data_idxs = np.delete(data_idxs, np.isin(data_idxs, slice_))
|
|
264
|
+
|
|
265
|
+
if return_idxs:
|
|
266
|
+
return tuple(loaders), tuple(cast(Subset, loader.dataset).indices for loader in loaders)
|
|
267
|
+
return tuple(loaders)
|
|
268
|
+
|
|
269
|
+
|
|
270
|
+
__all__ = ['BaseDataset', 'loader_init']
|