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.
Files changed (43) hide show
  1. {netloader-3.11.2 → netloader-3.12.0}/PKG-INFO +1 -1
  2. {netloader-3.11.2 → netloader-3.12.0}/netloader/__init__.py +1 -1
  3. {netloader-3.11.2 → netloader-3.12.0}/netloader/architectures/base.py +13 -1
  4. {netloader-3.11.2 → netloader-3.12.0}/netloader/architectures/encoder_decoder.py +15 -1
  5. {netloader-3.11.2 → netloader-3.12.0}/netloader/architectures/flows.py +13 -5
  6. netloader-3.12.0/netloader/data/__init__.py +8 -0
  7. netloader-3.12.0/netloader/data/datasets.py +270 -0
  8. netloader-3.11.2/netloader/data.py → netloader-3.12.0/netloader/data/structures.py +219 -310
  9. {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/base.py +5 -0
  10. {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/misc.py +5 -5
  11. {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/multi_layer.py +1 -1
  12. {netloader-3.11.2 → netloader-3.12.0}/netloader/models/misc.py +27 -15
  13. {netloader-3.11.2 → netloader-3.12.0}/netloader/network.py +91 -14
  14. {netloader-3.11.2 → netloader-3.12.0}/netloader/transforms.py +23 -5
  15. {netloader-3.11.2 → netloader-3.12.0}/netloader/utils/utils.py +36 -7
  16. {netloader-3.11.2 → netloader-3.12.0}/netloader.egg-info/PKG-INFO +1 -1
  17. {netloader-3.11.2 → netloader-3.12.0}/netloader.egg-info/SOURCES.txt +3 -1
  18. {netloader-3.11.2 → netloader-3.12.0}/LICENSE.txt +0 -0
  19. {netloader-3.11.2 → netloader-3.12.0}/README.md +0 -0
  20. {netloader-3.11.2 → netloader-3.12.0}/netloader/architectures/__init__.py +0 -0
  21. {netloader-3.11.2 → netloader-3.12.0}/netloader/architectures/utils.py +0 -0
  22. {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/__init__.py +0 -0
  23. {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/blocks.py +0 -0
  24. {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/convolutional.py +0 -0
  25. {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/flows.py +0 -0
  26. {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/linear.py +0 -0
  27. {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/pooling.py +0 -0
  28. {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/recurrent.py +0 -0
  29. {netloader-3.11.2 → netloader-3.12.0}/netloader/layers/utils.py +0 -0
  30. {netloader-3.11.2 → netloader-3.12.0}/netloader/loss_funcs.py +0 -0
  31. {netloader-3.11.2 → netloader-3.12.0}/netloader/models/__init__.py +0 -0
  32. {netloader-3.11.2 → netloader-3.12.0}/netloader/models/convnext.py +0 -0
  33. {netloader-3.11.2 → netloader-3.12.0}/netloader/networks/__init__.py +0 -0
  34. {netloader-3.11.2 → netloader-3.12.0}/netloader/schedulers.py +0 -0
  35. {netloader-3.11.2 → netloader-3.12.0}/netloader/utils/__init__.py +0 -0
  36. {netloader-3.11.2 → netloader-3.12.0}/netloader/utils/configs.py +0 -0
  37. {netloader-3.11.2 → netloader-3.12.0}/netloader/utils/transforms.py +0 -0
  38. {netloader-3.11.2 → netloader-3.12.0}/netloader/utils/types.py +0 -0
  39. {netloader-3.11.2 → netloader-3.12.0}/netloader.egg-info/dependency_links.txt +0 -0
  40. {netloader-3.11.2 → netloader-3.12.0}/netloader.egg-info/requires.txt +0 -0
  41. {netloader-3.11.2 → netloader-3.12.0}/netloader.egg-info/top_level.txt +0 -0
  42. {netloader-3.11.2 → netloader-3.12.0}/pyproject.toml +0 -0
  43. {netloader-3.11.2 → netloader-3.12.0}/setup.cfg +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: netloader
3
- Version: 3.11.2
3
+ Version: 3.12.0
4
4
  Summary: Utility to generate and train PyTorch neural network objects from JSON files
5
5
  Author-email: Ethan Tregidga <ethan.tregidga@epfl.ch>
6
6
  License-Expression: MIT
@@ -6,7 +6,7 @@ import logging
6
6
  import warnings
7
7
 
8
8
 
9
- __version__ = '3.11.2'
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, nn.ModuleList)
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']