netloader 3.11.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.
- netloader/__init__.py +61 -0
- netloader/architectures/__init__.py +48 -0
- netloader/architectures/base.py +1350 -0
- netloader/architectures/encoder_decoder.py +630 -0
- netloader/architectures/flows.py +501 -0
- netloader/architectures/utils.py +179 -0
- netloader/data.py +1092 -0
- netloader/layers/__init__.py +88 -0
- netloader/layers/base.py +367 -0
- netloader/layers/blocks.py +145 -0
- netloader/layers/convolutional.py +1008 -0
- netloader/layers/flows.py +123 -0
- netloader/layers/linear.py +445 -0
- netloader/layers/misc.py +879 -0
- netloader/layers/multi_layer.py +470 -0
- netloader/layers/pooling.py +298 -0
- netloader/layers/recurrent.py +163 -0
- netloader/layers/utils.py +168 -0
- netloader/loss_funcs.py +201 -0
- netloader/models/__init__.py +22 -0
- netloader/models/convnext.py +582 -0
- netloader/models/misc.py +209 -0
- netloader/network.py +895 -0
- netloader/networks/__init__.py +48 -0
- netloader/schedulers.py +409 -0
- netloader/transforms.py +847 -0
- netloader/utils/__init__.py +46 -0
- netloader/utils/configs.py +241 -0
- netloader/utils/transforms.py +22 -0
- netloader/utils/types.py +130 -0
- netloader/utils/utils.py +677 -0
- netloader-3.11.0.dist-info/METADATA +59 -0
- netloader-3.11.0.dist-info/RECORD +36 -0
- netloader-3.11.0.dist-info/WHEEL +5 -0
- netloader-3.11.0.dist-info/licenses/LICENSE.txt +21 -0
- netloader-3.11.0.dist-info/top_level.txt +1 -0
netloader/data.py
ADDED
|
@@ -0,0 +1,1092 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Base dataset classes for use with BaseArchitecture
|
|
3
|
+
"""
|
|
4
|
+
from __future__ import annotations
|
|
5
|
+
import logging as log
|
|
6
|
+
from types import ModuleType
|
|
7
|
+
from itertools import repeat
|
|
8
|
+
from typing import (
|
|
9
|
+
Any,
|
|
10
|
+
Self,
|
|
11
|
+
Generic,
|
|
12
|
+
Literal,
|
|
13
|
+
Sequence,
|
|
14
|
+
Iterator,
|
|
15
|
+
Protocol,
|
|
16
|
+
TypeVar,
|
|
17
|
+
cast,
|
|
18
|
+
overload,
|
|
19
|
+
)
|
|
20
|
+
|
|
21
|
+
import torch
|
|
22
|
+
import numpy as np
|
|
23
|
+
from torch import Tensor
|
|
24
|
+
from torch.utils.data import Dataset, DataLoader, Subset
|
|
25
|
+
from numpy import ndarray
|
|
26
|
+
|
|
27
|
+
from netloader.utils.types import (
|
|
28
|
+
DataLike,
|
|
29
|
+
ArrayLike,
|
|
30
|
+
DataT,
|
|
31
|
+
ArrayT,
|
|
32
|
+
DatasetT,
|
|
33
|
+
DataListT,
|
|
34
|
+
ArrayCT,
|
|
35
|
+
P,
|
|
36
|
+
)
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
InArrayT_contra = TypeVar('InArrayT_contra', bound=DataLike, contravariant=True)
|
|
40
|
+
OutArrayT_co = TypeVar('OutArrayT_co', bound=DataLike, covariant=True)
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
class ApplyFunc(Protocol[InArrayT_contra, P, OutArrayT_co]):
|
|
44
|
+
"""
|
|
45
|
+
Protocol for functions that can be applied to Data objects in the apply method of Data and
|
|
46
|
+
DataList.
|
|
47
|
+
"""
|
|
48
|
+
def __call__(self, data: InArrayT_contra, /, *args: P.args, **kwargs: P.kwargs) -> OutArrayT_co:
|
|
49
|
+
"""
|
|
50
|
+
Parameters
|
|
51
|
+
----------
|
|
52
|
+
data : DataT
|
|
53
|
+
Data to apply the function to
|
|
54
|
+
*args
|
|
55
|
+
Optional arguments to pass to the function
|
|
56
|
+
**kwargs
|
|
57
|
+
Optional keyword arguments to pass to the function
|
|
58
|
+
|
|
59
|
+
Returns
|
|
60
|
+
-------
|
|
61
|
+
DataT
|
|
62
|
+
Result of applying the function to the data
|
|
63
|
+
"""
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
class Data(Generic[ArrayCT]):
|
|
67
|
+
"""
|
|
68
|
+
Stores data with uncertainties for uncertainty handling in BaseArchitecture.
|
|
69
|
+
|
|
70
|
+
Attributes
|
|
71
|
+
----------
|
|
72
|
+
data : ArrayCT
|
|
73
|
+
Data
|
|
74
|
+
uncertainty : ArrayCT | None
|
|
75
|
+
Data uncertainty
|
|
76
|
+
"""
|
|
77
|
+
def __init__(self, data: ArrayCT, uncertainty: ArrayCT | None = None) -> None:
|
|
78
|
+
"""
|
|
79
|
+
Parameters
|
|
80
|
+
----------
|
|
81
|
+
data : ArrayCT
|
|
82
|
+
Data
|
|
83
|
+
uncertainty : ArrayCT | None, Optional
|
|
84
|
+
Data uncertainty
|
|
85
|
+
"""
|
|
86
|
+
self.shape: tuple[int, ...] = tuple(data.shape)
|
|
87
|
+
self.data: ArrayCT = data
|
|
88
|
+
self.uncertainty: ArrayCT | None = uncertainty
|
|
89
|
+
|
|
90
|
+
def __len__(self) -> int:
|
|
91
|
+
"""
|
|
92
|
+
Returns the number of samples in the data.
|
|
93
|
+
|
|
94
|
+
Returns
|
|
95
|
+
-------
|
|
96
|
+
int
|
|
97
|
+
Number of samples in the data
|
|
98
|
+
"""
|
|
99
|
+
return len(self.data)
|
|
100
|
+
|
|
101
|
+
def __getitem__(self, idx: int | slice) -> Data[ArrayCT]:
|
|
102
|
+
"""
|
|
103
|
+
Gets a subset of the data and uncertainty.
|
|
104
|
+
|
|
105
|
+
Parameters
|
|
106
|
+
----------
|
|
107
|
+
idx : int | slice
|
|
108
|
+
Index or slice to get
|
|
109
|
+
|
|
110
|
+
Returns
|
|
111
|
+
-------
|
|
112
|
+
Data[ArrayCT]
|
|
113
|
+
Data with subset of data and uncertainty
|
|
114
|
+
"""
|
|
115
|
+
return Data(
|
|
116
|
+
self.data[idx],
|
|
117
|
+
uncertainty=self.uncertainty[idx] if self.uncertainty is not None else None
|
|
118
|
+
)
|
|
119
|
+
|
|
120
|
+
def __repr__(self) -> str:
|
|
121
|
+
return (f'Data(shape={self.shape}, type={self.data.dtype}, '
|
|
122
|
+
f'uncertainty={self.uncertainty is not None})')
|
|
123
|
+
|
|
124
|
+
@staticmethod
|
|
125
|
+
def _apply(
|
|
126
|
+
data: ArrayCT,
|
|
127
|
+
func: str | ApplyFunc[ArrayCT, ..., ArrayCT],
|
|
128
|
+
*args: Any,
|
|
129
|
+
types: tuple[type[DataLike], ...] = (Tensor, ndarray),
|
|
130
|
+
**kwargs: Any) -> ArrayCT:
|
|
131
|
+
if data not in types:
|
|
132
|
+
return data
|
|
133
|
+
if isinstance(func, str):
|
|
134
|
+
return getattr(data, func)(*args, **kwargs)
|
|
135
|
+
return func(data, *args, **kwargs)
|
|
136
|
+
|
|
137
|
+
@staticmethod
|
|
138
|
+
@overload
|
|
139
|
+
def collate(data: list[Data[ArrayCT]], *, data_field: Literal[True]) -> Data[ArrayCT]: ...
|
|
140
|
+
|
|
141
|
+
@staticmethod
|
|
142
|
+
@overload
|
|
143
|
+
def collate(data: list[Data[ArrayCT]], *, data_field: Literal[False]) -> ArrayCT: ...
|
|
144
|
+
|
|
145
|
+
@staticmethod
|
|
146
|
+
@overload
|
|
147
|
+
def collate(data: list[Data[ArrayCT]], *, data_field: bool) -> ArrayCT | Data[ArrayCT]: ...
|
|
148
|
+
|
|
149
|
+
@staticmethod
|
|
150
|
+
def collate(
|
|
151
|
+
data: list[Data[ArrayCT]],
|
|
152
|
+
*,
|
|
153
|
+
data_field: bool = True) -> ArrayCT | Data[ArrayCT]:
|
|
154
|
+
"""
|
|
155
|
+
Collates a list of Data objects into a single Data object.
|
|
156
|
+
|
|
157
|
+
Parameters
|
|
158
|
+
----------
|
|
159
|
+
data : list[Data[ArrayCT]]
|
|
160
|
+
List of Data objects to collate
|
|
161
|
+
data_field : bool, Optional
|
|
162
|
+
If True, returns a Data object, else returns the collated data as an ArrayLike,
|
|
163
|
+
default = True
|
|
164
|
+
|
|
165
|
+
Returns
|
|
166
|
+
-------
|
|
167
|
+
ArrayCT | Data[ArrayCT]
|
|
168
|
+
Collated ArrayLike or Data object
|
|
169
|
+
"""
|
|
170
|
+
module: ModuleType = torch if isinstance(data[0].data, Tensor) else np
|
|
171
|
+
new_data: ArrayCT
|
|
172
|
+
datum: Data[Any]
|
|
173
|
+
|
|
174
|
+
if data[0].uncertainty is None:
|
|
175
|
+
new_data = module.concat([datum.data for datum in data])
|
|
176
|
+
return Data(new_data) if data_field else new_data
|
|
177
|
+
|
|
178
|
+
new_data = module.concat([datum.concat() for datum in data])
|
|
179
|
+
return Data(*new_data.swapaxes(0, 1)) if data_field else new_data
|
|
180
|
+
|
|
181
|
+
def apply(
|
|
182
|
+
self,
|
|
183
|
+
func: str | ApplyFunc[ArrayCT, ..., ArrayCT],
|
|
184
|
+
*args: Any,
|
|
185
|
+
types: tuple[type[DataLike], ...] = (Tensor, ndarray),
|
|
186
|
+
**kwargs: Any) -> Self:
|
|
187
|
+
"""
|
|
188
|
+
Applies a function or method to the data and uncertainty.
|
|
189
|
+
|
|
190
|
+
Parameters
|
|
191
|
+
----------
|
|
192
|
+
func : str | ApplyFunc[ArrayCT]
|
|
193
|
+
Function, or method if string, to apply to the data and uncertainty
|
|
194
|
+
*args
|
|
195
|
+
Arguments to pass to the function or method
|
|
196
|
+
types : tuple[type[ArrayLike], ...]
|
|
197
|
+
Types to apply the function or method to, if data and uncertainty is not an instance of
|
|
198
|
+
types, then data and uncertainty are returned unchanged, default = (Tensor, ndarray)
|
|
199
|
+
**kwargs
|
|
200
|
+
Keyword arguments to pass to the function or method
|
|
201
|
+
|
|
202
|
+
Returns
|
|
203
|
+
-------
|
|
204
|
+
Self
|
|
205
|
+
Self with function or method applied to the data and uncertainty
|
|
206
|
+
"""
|
|
207
|
+
if not isinstance(self.data, types):
|
|
208
|
+
return self
|
|
209
|
+
|
|
210
|
+
self.data = self._apply(self.data, func, *args, types=types, **kwargs)
|
|
211
|
+
self.uncertainty = None if self.uncertainty is None else self._apply(
|
|
212
|
+
self.uncertainty,
|
|
213
|
+
func,
|
|
214
|
+
*args,
|
|
215
|
+
types=types,
|
|
216
|
+
**kwargs,
|
|
217
|
+
)
|
|
218
|
+
return self
|
|
219
|
+
|
|
220
|
+
def clone(self) -> Data[ArrayCT]:
|
|
221
|
+
"""
|
|
222
|
+
Clones the data and uncertainty.
|
|
223
|
+
|
|
224
|
+
Identical to Data.copy().
|
|
225
|
+
|
|
226
|
+
Returns
|
|
227
|
+
-------
|
|
228
|
+
Data[ArrayCT]
|
|
229
|
+
Cloned Data
|
|
230
|
+
"""
|
|
231
|
+
data: ArrayCT
|
|
232
|
+
uncertainty: ArrayCT | None
|
|
233
|
+
|
|
234
|
+
if isinstance(self.data, ndarray):
|
|
235
|
+
data = self.data.copy()
|
|
236
|
+
uncertainty = self.uncertainty.copy() if self.uncertainty is not None else None
|
|
237
|
+
else:
|
|
238
|
+
data = self.data.clone()
|
|
239
|
+
uncertainty = self.uncertainty.clone() if self.uncertainty is not None else None
|
|
240
|
+
return Data(data, uncertainty=uncertainty)
|
|
241
|
+
|
|
242
|
+
def concat(self, dim: int = 0) -> ArrayCT:
|
|
243
|
+
"""
|
|
244
|
+
Concatenates data and uncertainty for passing into a network.
|
|
245
|
+
|
|
246
|
+
Parameters
|
|
247
|
+
----------
|
|
248
|
+
dim : int, Optional
|
|
249
|
+
Dimension to concatenate along, default = 0
|
|
250
|
+
|
|
251
|
+
Returns
|
|
252
|
+
-------
|
|
253
|
+
ArrayCT
|
|
254
|
+
Data and uncertainty concatenated along the specified dimension if uncertainty is not
|
|
255
|
+
None, else just data
|
|
256
|
+
"""
|
|
257
|
+
module: ModuleType = torch if isinstance(self.data, Tensor) else np
|
|
258
|
+
kwargs: dict[str, int] = {'dim' if isinstance(self.data, Tensor) else 'axis': dim}
|
|
259
|
+
|
|
260
|
+
if self.uncertainty is not None:
|
|
261
|
+
return module.concat((self.data, self.uncertainty), **kwargs)
|
|
262
|
+
return self.data
|
|
263
|
+
|
|
264
|
+
copy = clone
|
|
265
|
+
|
|
266
|
+
def cpu(self) -> Self:
|
|
267
|
+
"""
|
|
268
|
+
Moves data and uncertainty to CPU if they are Tensors.
|
|
269
|
+
|
|
270
|
+
Returns
|
|
271
|
+
-------
|
|
272
|
+
Self
|
|
273
|
+
Self with data and uncertainty on CPU
|
|
274
|
+
"""
|
|
275
|
+
if isinstance(self.data, Tensor):
|
|
276
|
+
self.data = self.data.cpu()
|
|
277
|
+
|
|
278
|
+
if isinstance(self.uncertainty, Tensor):
|
|
279
|
+
self.uncertainty = self.uncertainty.cpu()
|
|
280
|
+
return self.apply('cpu')
|
|
281
|
+
|
|
282
|
+
def detach(self) -> Self:
|
|
283
|
+
"""
|
|
284
|
+
Detaches data and uncertainty from the computation graph if they are Tensors.
|
|
285
|
+
|
|
286
|
+
Returns
|
|
287
|
+
-------
|
|
288
|
+
Self
|
|
289
|
+
Self with data and uncertainty detached from the computation graph
|
|
290
|
+
"""
|
|
291
|
+
if isinstance(self.data, Tensor):
|
|
292
|
+
self.data = self.data.detach()
|
|
293
|
+
|
|
294
|
+
if isinstance(self.uncertainty, Tensor):
|
|
295
|
+
self.uncertainty = self.uncertainty.detach()
|
|
296
|
+
return self
|
|
297
|
+
|
|
298
|
+
def numpy(self) -> Data[ndarray]:
|
|
299
|
+
"""
|
|
300
|
+
Converts data and uncertainty to numpy arrays if they are Tensors.
|
|
301
|
+
|
|
302
|
+
Returns
|
|
303
|
+
-------
|
|
304
|
+
Data[ndarray]
|
|
305
|
+
New Data object with data and uncertainty as numpy arrays
|
|
306
|
+
"""
|
|
307
|
+
return Data(
|
|
308
|
+
self.data.cpu().numpy() if isinstance(self.data, Tensor) else self.data,
|
|
309
|
+
uncertainty=self.uncertainty.cpu().numpy() if isinstance(self.uncertainty, Tensor) else
|
|
310
|
+
self.uncertainty
|
|
311
|
+
)
|
|
312
|
+
|
|
313
|
+
def tensor(self) -> Data[Tensor]:
|
|
314
|
+
"""
|
|
315
|
+
Converts data and uncertainty to tensors if they are numpy arrays.
|
|
316
|
+
|
|
317
|
+
Returns
|
|
318
|
+
-------
|
|
319
|
+
Data[Tensor]
|
|
320
|
+
New Data object with data and uncertainty as tensors
|
|
321
|
+
"""
|
|
322
|
+
return Data(
|
|
323
|
+
torch.from_numpy(self.data) if isinstance(self.data, ndarray) else self.data,
|
|
324
|
+
uncertainty=torch.from_numpy(self.uncertainty) if isinstance(self.uncertainty, ndarray)
|
|
325
|
+
else self.uncertainty
|
|
326
|
+
)
|
|
327
|
+
|
|
328
|
+
def to(self, *args: Any, **kwargs: Any) -> Self:
|
|
329
|
+
"""
|
|
330
|
+
Move and/or cast the parameters and buffers.
|
|
331
|
+
|
|
332
|
+
Parameters
|
|
333
|
+
----------
|
|
334
|
+
*args
|
|
335
|
+
Arguments to pass to the to method of data and uncertainty
|
|
336
|
+
**kwargs
|
|
337
|
+
Keyword arguments to pass to the to method of data and uncertainty
|
|
338
|
+
|
|
339
|
+
Returns
|
|
340
|
+
-------
|
|
341
|
+
Self
|
|
342
|
+
Self with data and uncertainty moved and/or cast
|
|
343
|
+
"""
|
|
344
|
+
if isinstance(self.data, Tensor):
|
|
345
|
+
self.data = self.data.to(*args, **kwargs)
|
|
346
|
+
|
|
347
|
+
if isinstance(self.uncertainty, Tensor):
|
|
348
|
+
self.uncertainty = self.uncertainty.to(*args, **kwargs)
|
|
349
|
+
return self
|
|
350
|
+
|
|
351
|
+
|
|
352
|
+
class DataList(Generic[DataT]):
|
|
353
|
+
"""
|
|
354
|
+
A list that stores tensors, arrays, or Datas and provides batch operations.
|
|
355
|
+
"""
|
|
356
|
+
def __init__(self, data: list[DataT]) -> None:
|
|
357
|
+
"""
|
|
358
|
+
Parameters
|
|
359
|
+
----------
|
|
360
|
+
data : list[DataT]
|
|
361
|
+
List of Data, Tensor, or ndarray objects
|
|
362
|
+
"""
|
|
363
|
+
self._data: list[DataT] = data
|
|
364
|
+
|
|
365
|
+
if any(len(data[0]) != len(datum) for datum in data[1:]):
|
|
366
|
+
raise ValueError(f'All elements in DataList must have the same length, got lengths: '
|
|
367
|
+
f'{[len(datum) for datum in data]}')
|
|
368
|
+
|
|
369
|
+
def __add__(self, other: DataT | DataList[DataT]) -> DataList[DataT]:
|
|
370
|
+
"""
|
|
371
|
+
Adds another DataList or data to the DataList.
|
|
372
|
+
|
|
373
|
+
Parameters
|
|
374
|
+
----------
|
|
375
|
+
other : DataT | DataList[DataT]
|
|
376
|
+
DataList or single element to add to the DataList
|
|
377
|
+
|
|
378
|
+
Returns
|
|
379
|
+
-------
|
|
380
|
+
DataList[DataT]
|
|
381
|
+
New DataList with the other DataList or data added to the end of the DataList
|
|
382
|
+
"""
|
|
383
|
+
if len(self) != len(other):
|
|
384
|
+
raise ValueError(f'Lengths of DataList ({len(self)}) and {type(other).__name__} '
|
|
385
|
+
f'({len(other)}) must be the same')
|
|
386
|
+
return DataList(self._data + (list(other) if isinstance(other, DataList) else [other]))
|
|
387
|
+
|
|
388
|
+
def __getitem__(self, idx: int | slice) -> DataList[DataT]:
|
|
389
|
+
"""
|
|
390
|
+
Gets a subset of each element in the DataList.
|
|
391
|
+
|
|
392
|
+
Parameters
|
|
393
|
+
----------
|
|
394
|
+
idx: int | slice
|
|
395
|
+
Index or slice to get
|
|
396
|
+
|
|
397
|
+
Returns
|
|
398
|
+
-------
|
|
399
|
+
DataList[DataT]
|
|
400
|
+
DataList with subset of each element in the DataList
|
|
401
|
+
"""
|
|
402
|
+
return self.get(idx, list_=False)
|
|
403
|
+
|
|
404
|
+
def __iter__(self) -> Iterator[DataT]:
|
|
405
|
+
"""
|
|
406
|
+
Iterates over each element in the DataList.
|
|
407
|
+
|
|
408
|
+
Returns
|
|
409
|
+
-------
|
|
410
|
+
Iterator[DataT]
|
|
411
|
+
Iterator over each element in the DataList
|
|
412
|
+
"""
|
|
413
|
+
return self.iter(list_=True)
|
|
414
|
+
|
|
415
|
+
def __len__(self) -> int:
|
|
416
|
+
"""
|
|
417
|
+
Returns the number of samples in each element.
|
|
418
|
+
|
|
419
|
+
Returns
|
|
420
|
+
-------
|
|
421
|
+
int
|
|
422
|
+
Number of samples in each element
|
|
423
|
+
"""
|
|
424
|
+
return self.len(list_=False)
|
|
425
|
+
|
|
426
|
+
def __repr__(self) -> str:
|
|
427
|
+
"""
|
|
428
|
+
Returns a string representation of the DataList.
|
|
429
|
+
|
|
430
|
+
Returns
|
|
431
|
+
-------
|
|
432
|
+
str
|
|
433
|
+
String representation of the DataList
|
|
434
|
+
"""
|
|
435
|
+
return (f'DataList(shapes={[tuple(datum.shape) for datum in self._data]}, '
|
|
436
|
+
f'types={[datum.__class__.__name__ for datum in self._data]})')
|
|
437
|
+
|
|
438
|
+
@staticmethod
|
|
439
|
+
def _apply(
|
|
440
|
+
data: DataT,
|
|
441
|
+
func: str | ApplyFunc[DataT, ..., DataT],
|
|
442
|
+
*args: Any,
|
|
443
|
+
types: tuple[type[DataLike], ...] = (ndarray, Tensor, Data),
|
|
444
|
+
**kwargs: Any) -> DataT:
|
|
445
|
+
"""
|
|
446
|
+
Applies a function or method to the data if it is an instance of types, else returns the
|
|
447
|
+
data unchanged.
|
|
448
|
+
|
|
449
|
+
Parameters
|
|
450
|
+
----------
|
|
451
|
+
data : DataT
|
|
452
|
+
Data to apply function to
|
|
453
|
+
func : str | ApplyFunc[ArrayT]
|
|
454
|
+
Function, or method if string, to apply to the data
|
|
455
|
+
*args
|
|
456
|
+
Arguments to pass to the function or method
|
|
457
|
+
types : tuple[type[DataLike], ...]
|
|
458
|
+
Types to apply the function or method to, if data is not an instance of types, then
|
|
459
|
+
data is returned unchanged, default = (ndarray, Tensor, Data)
|
|
460
|
+
**kwargs
|
|
461
|
+
Keyword arguments to pass to the function or method
|
|
462
|
+
|
|
463
|
+
Returns
|
|
464
|
+
-------
|
|
465
|
+
DataT
|
|
466
|
+
Result of applying the function or method to the data if it is an instance of types,
|
|
467
|
+
else returns the data unchanged
|
|
468
|
+
"""
|
|
469
|
+
if not isinstance(data, types):
|
|
470
|
+
return data
|
|
471
|
+
if isinstance(func, str):
|
|
472
|
+
return getattr(data, func)(*args, **kwargs)
|
|
473
|
+
return func(data, *args, **kwargs)
|
|
474
|
+
|
|
475
|
+
@staticmethod
|
|
476
|
+
@overload
|
|
477
|
+
def collate(
|
|
478
|
+
data: list[DataList[DataT]],
|
|
479
|
+
*,
|
|
480
|
+
data_field: Literal[True]) -> DataList[DataT]: ...
|
|
481
|
+
|
|
482
|
+
@staticmethod
|
|
483
|
+
@overload
|
|
484
|
+
def collate(
|
|
485
|
+
data: list[DataList[DataT]],
|
|
486
|
+
*,
|
|
487
|
+
data_field: Literal[False]) -> DataList[ArrayT]: ...
|
|
488
|
+
|
|
489
|
+
@staticmethod
|
|
490
|
+
@overload
|
|
491
|
+
def collate(
|
|
492
|
+
data: list[DataList[DataT]],
|
|
493
|
+
*,
|
|
494
|
+
data_field: bool) -> DataList[ArrayT] | DataList[DataT]: ...
|
|
495
|
+
|
|
496
|
+
@staticmethod
|
|
497
|
+
def collate(
|
|
498
|
+
data: list[DataList[DataT]],
|
|
499
|
+
*,
|
|
500
|
+
data_field: bool = True) -> DataList[ArrayT] | DataList[DataT]:
|
|
501
|
+
"""
|
|
502
|
+
Collates a list of DataList objects into a single DataList object.
|
|
503
|
+
|
|
504
|
+
Parameters
|
|
505
|
+
----------
|
|
506
|
+
data : list[DataList[DataT]]
|
|
507
|
+
List of DataList objects to collate
|
|
508
|
+
data_field : bool, Optional
|
|
509
|
+
If True, collates Data elements into Data object, else collates into ArrayLike,
|
|
510
|
+
default = True
|
|
511
|
+
|
|
512
|
+
Returns
|
|
513
|
+
-------
|
|
514
|
+
DataList[ArrayT] | DataList[DataT]
|
|
515
|
+
Collated DataList object
|
|
516
|
+
"""
|
|
517
|
+
i: int
|
|
518
|
+
element_data: list[DataLike]
|
|
519
|
+
new_data: list[DataLike] = []
|
|
520
|
+
datum: DataList[DataT]
|
|
521
|
+
|
|
522
|
+
for i in range(data[0].len(True)):
|
|
523
|
+
element_data = [datum.get(i, list_=True) for datum in data]
|
|
524
|
+
new_data.append(
|
|
525
|
+
Data.collate(cast(list[Data], element_data), data_field=data_field)
|
|
526
|
+
if isinstance(element_data[0], Data) else
|
|
527
|
+
torch.concat(cast(list[Tensor], element_data))
|
|
528
|
+
if isinstance(element_data[0], Tensor) else
|
|
529
|
+
np.concat(cast(list[ndarray], element_data)),
|
|
530
|
+
)
|
|
531
|
+
return cast(DataList[ArrayT] | DataList[DataT], DataList(new_data))
|
|
532
|
+
|
|
533
|
+
def append(self, data: DataT) -> None:
|
|
534
|
+
"""
|
|
535
|
+
Appends an element to the DataList.
|
|
536
|
+
|
|
537
|
+
Parameters
|
|
538
|
+
----------
|
|
539
|
+
data : DataT
|
|
540
|
+
Element to append to the DataList
|
|
541
|
+
"""
|
|
542
|
+
if self.len(list_=False) != len(data):
|
|
543
|
+
raise ValueError(f'Element to append must have the same length as the DataList, got '
|
|
544
|
+
f'lengths: {self.len(list_=False)} and {len(data)}')
|
|
545
|
+
self._data.append(data)
|
|
546
|
+
|
|
547
|
+
def apply(
|
|
548
|
+
self,
|
|
549
|
+
func: str | list[str | ApplyFunc[DataT, ..., DataT]] | ApplyFunc[DataT, ..., DataT],
|
|
550
|
+
*args: Any,
|
|
551
|
+
types: tuple[type[DataLike], ...] = (Data, Tensor, ndarray),
|
|
552
|
+
**kwargs: Any) -> Self:
|
|
553
|
+
"""
|
|
554
|
+
Applies a function or method to each element in the DataList.
|
|
555
|
+
|
|
556
|
+
Parameters
|
|
557
|
+
----------
|
|
558
|
+
func : str | list[str | ApplyFunc[ArrayT]] | ApplyFunc[DataT]
|
|
559
|
+
Function(s), or method(s) if string, to apply to each element in the DataList
|
|
560
|
+
*args
|
|
561
|
+
Arguments to pass to the function or method
|
|
562
|
+
types : tuple[type[DataLike], ...]
|
|
563
|
+
Types to apply the function or method to, if an element is not an instance of types,
|
|
564
|
+
then that element is returned unchanged, default = (Data, Tensor, ndarray)
|
|
565
|
+
**kwargs
|
|
566
|
+
Keyword arguments to pass to the function or method
|
|
567
|
+
|
|
568
|
+
Returns
|
|
569
|
+
-------
|
|
570
|
+
Self
|
|
571
|
+
Self with function or method applied to each element in the DataList
|
|
572
|
+
"""
|
|
573
|
+
datum: DataT
|
|
574
|
+
func_: str | ApplyFunc[DataT, ..., DataT]
|
|
575
|
+
|
|
576
|
+
if isinstance(func, list) and len(func) != len(self._data):
|
|
577
|
+
raise ValueError(f'Length of function list ({len(func)}) must match length of DataList '
|
|
578
|
+
f'({len(self._data)})')
|
|
579
|
+
|
|
580
|
+
self._data = [
|
|
581
|
+
self._apply(datum, func_, *args, types=types, **kwargs)
|
|
582
|
+
for datum, func_ in zip(self._data, func if isinstance(func, list) else repeat(func))
|
|
583
|
+
]
|
|
584
|
+
return self
|
|
585
|
+
|
|
586
|
+
def clone(self) -> DataList[DataT]:
|
|
587
|
+
"""
|
|
588
|
+
Clones the DataList.
|
|
589
|
+
|
|
590
|
+
Identical to DataList.copy().
|
|
591
|
+
|
|
592
|
+
Returns
|
|
593
|
+
-------
|
|
594
|
+
DataList[DataT]
|
|
595
|
+
Cloned DataList
|
|
596
|
+
"""
|
|
597
|
+
data: DataT
|
|
598
|
+
return DataList([
|
|
599
|
+
cast(DataT, data.clone() if isinstance(data, Tensor) else data.copy())
|
|
600
|
+
for data in self._data
|
|
601
|
+
])
|
|
602
|
+
|
|
603
|
+
copy = clone
|
|
604
|
+
|
|
605
|
+
def cpu(self) -> Self:
|
|
606
|
+
"""
|
|
607
|
+
Moves all tensors to CPU
|
|
608
|
+
|
|
609
|
+
Returns
|
|
610
|
+
-------
|
|
611
|
+
Self
|
|
612
|
+
Self with all tensors moved to CPU
|
|
613
|
+
"""
|
|
614
|
+
return self.apply('cpu', types=(Tensor, Data))
|
|
615
|
+
|
|
616
|
+
def detach(self) -> Self:
|
|
617
|
+
"""
|
|
618
|
+
Detaches all tensors from the computation graph
|
|
619
|
+
|
|
620
|
+
Returns
|
|
621
|
+
-------
|
|
622
|
+
Self
|
|
623
|
+
Self with all tensors detached from the computation graph
|
|
624
|
+
"""
|
|
625
|
+
return self.apply('detach', types=(Tensor, Data))
|
|
626
|
+
|
|
627
|
+
def extend(self, data: Sequence[DataT] | DataList[DataT]) -> None:
|
|
628
|
+
"""
|
|
629
|
+
Extends the DataList by appending elements from the iterable.
|
|
630
|
+
|
|
631
|
+
Parameters
|
|
632
|
+
----------
|
|
633
|
+
data : Sequence[DataT] | DataList[DataT]
|
|
634
|
+
Elements to extend the DataList with
|
|
635
|
+
"""
|
|
636
|
+
datum: DataT
|
|
637
|
+
|
|
638
|
+
for datum in data:
|
|
639
|
+
self.append(datum)
|
|
640
|
+
|
|
641
|
+
@overload
|
|
642
|
+
def get(self, idx: int, list_: Literal[True]) -> DataT: ...
|
|
643
|
+
|
|
644
|
+
@overload
|
|
645
|
+
def get(self, idx: slice, list_: Literal[True]) -> list[DataT]: ...
|
|
646
|
+
|
|
647
|
+
@overload
|
|
648
|
+
def get(self, idx: int | slice, list_: Literal[False]) -> DataList[DataT]: ...
|
|
649
|
+
|
|
650
|
+
@overload
|
|
651
|
+
def get(self, idx: int, list_: bool) -> DataT | DataList[DataT]: ...
|
|
652
|
+
|
|
653
|
+
@overload
|
|
654
|
+
def get(self, idx: slice, list_: bool) -> list[DataT] | DataList[DataT]: ...
|
|
655
|
+
|
|
656
|
+
def get(self, idx: int | slice, list_: bool = False) -> DataT | list[DataT] | DataList[DataT]:
|
|
657
|
+
"""
|
|
658
|
+
Gets a subset of the DataList or a subset of each element in the DataList.
|
|
659
|
+
|
|
660
|
+
Parameters
|
|
661
|
+
----------
|
|
662
|
+
idx : int | slice
|
|
663
|
+
Index or slice to get
|
|
664
|
+
list_ : bool, Optional
|
|
665
|
+
If True, returns a subset of the DataList, else returns a subset of each element in the
|
|
666
|
+
DataList, default = False
|
|
667
|
+
|
|
668
|
+
Returns
|
|
669
|
+
-------
|
|
670
|
+
DataT | list[DataT] | DataList[DataT]
|
|
671
|
+
Subset of the DataList or subset of each element in the DataList
|
|
672
|
+
"""
|
|
673
|
+
data: DataT
|
|
674
|
+
|
|
675
|
+
if list_:
|
|
676
|
+
return self._data[idx]
|
|
677
|
+
return DataList(cast(list[DataT], [data[idx] for data in self]))
|
|
678
|
+
|
|
679
|
+
def get_data(self) -> list[DataT]:
|
|
680
|
+
"""
|
|
681
|
+
Gets the underlying data list.
|
|
682
|
+
|
|
683
|
+
Returns
|
|
684
|
+
-------
|
|
685
|
+
list[DataT]
|
|
686
|
+
Underlying data list
|
|
687
|
+
"""
|
|
688
|
+
return self._data
|
|
689
|
+
|
|
690
|
+
def insert(self, idx: int, data: DataT) -> None:
|
|
691
|
+
"""
|
|
692
|
+
Inserts an element into the DataList at the specified index.
|
|
693
|
+
|
|
694
|
+
Parameters
|
|
695
|
+
----------
|
|
696
|
+
idx : int
|
|
697
|
+
Index to insert the element at
|
|
698
|
+
data : DataT
|
|
699
|
+
Element to insert into the DataList
|
|
700
|
+
"""
|
|
701
|
+
if self.len(list_=False) != len(data):
|
|
702
|
+
raise ValueError(f'Element to insert must have the same length as the DataList, got '
|
|
703
|
+
f'lengths: {self.len(list_=False)} and {len(data)}')
|
|
704
|
+
self._data.insert(idx, data)
|
|
705
|
+
|
|
706
|
+
@overload
|
|
707
|
+
def iter(self, list_: Literal[True]) -> Iterator[DataT]: ...
|
|
708
|
+
|
|
709
|
+
@overload
|
|
710
|
+
def iter(self, list_: Literal[False]) -> Iterator[DataList[DataT]]: ...
|
|
711
|
+
|
|
712
|
+
@overload
|
|
713
|
+
def iter(self, list_: bool) -> Iterator[DataT] | Iterator[DataList[DataT]]: ...
|
|
714
|
+
|
|
715
|
+
def iter(self, list_: bool = True) -> Iterator[DataT] | Iterator[DataList[DataT]]:
|
|
716
|
+
"""
|
|
717
|
+
Iterates over each element in the DataList.
|
|
718
|
+
|
|
719
|
+
Parameters
|
|
720
|
+
----------
|
|
721
|
+
list_ : bool, Optional
|
|
722
|
+
If True, iterates over the DataList, else iterates over each element in the DataList,
|
|
723
|
+
default = True
|
|
724
|
+
|
|
725
|
+
Returns
|
|
726
|
+
-------
|
|
727
|
+
Iterator[DataT] | Iterator[DataList[DataT]]
|
|
728
|
+
Iterator over each element in the DataList
|
|
729
|
+
"""
|
|
730
|
+
i: int
|
|
731
|
+
|
|
732
|
+
for i in range(self.len(True)):
|
|
733
|
+
yield self.get(i, list_=list_)
|
|
734
|
+
|
|
735
|
+
def len(self, list_: bool = True) -> int:
|
|
736
|
+
"""
|
|
737
|
+
Gets the length of the DataList or the length of each element in the DataList.
|
|
738
|
+
|
|
739
|
+
Parameters
|
|
740
|
+
----------
|
|
741
|
+
list_ : bool, Optional
|
|
742
|
+
If True, returns the length of the DataList, else returns the length of each element
|
|
743
|
+
in the DataList, default = True
|
|
744
|
+
|
|
745
|
+
Returns
|
|
746
|
+
-------
|
|
747
|
+
int
|
|
748
|
+
Length of the DataList or length of each element in the DataList
|
|
749
|
+
"""
|
|
750
|
+
return len(self._data) if list_ else len(self._data[0])
|
|
751
|
+
|
|
752
|
+
def numpy(self) -> DataList[ndarray | Data[ndarray]]:
|
|
753
|
+
"""
|
|
754
|
+
Converts all tensors to numpy arrays.
|
|
755
|
+
|
|
756
|
+
Returns
|
|
757
|
+
-------
|
|
758
|
+
DataList[ndarray | Data[ndarray]]
|
|
759
|
+
DataList with all tensors converted to numpy arrays
|
|
760
|
+
"""
|
|
761
|
+
data: DataT
|
|
762
|
+
return DataList(cast(list[ndarray | Data[ndarray]], [
|
|
763
|
+
data.cpu().numpy() if hasattr(data, 'cpu') else data for data in self
|
|
764
|
+
]))
|
|
765
|
+
|
|
766
|
+
def tensor(self) -> DataList[Tensor | Data[Tensor]]:
|
|
767
|
+
"""
|
|
768
|
+
Converts all numpy arrays to tensors.
|
|
769
|
+
|
|
770
|
+
Returns
|
|
771
|
+
-------
|
|
772
|
+
DataList[Tensor | Data[Tensor]]
|
|
773
|
+
DataList with all numpy arrays converted to tensors
|
|
774
|
+
"""
|
|
775
|
+
data: DataT
|
|
776
|
+
return DataList(cast(list[Tensor | Data[Tensor]], [
|
|
777
|
+
torch.from_numpy(data) if isinstance(data, ndarray) else
|
|
778
|
+
data.tensor() if isinstance(data, Data) else data for data in self
|
|
779
|
+
]))
|
|
780
|
+
|
|
781
|
+
def to(self, *args: Any, **kwargs: Any) -> Self:
|
|
782
|
+
"""
|
|
783
|
+
Move and/or cast the parameters and buffers.
|
|
784
|
+
|
|
785
|
+
Parameters
|
|
786
|
+
----------
|
|
787
|
+
*args
|
|
788
|
+
Arguments to pass to the to method of each element in the DataList
|
|
789
|
+
**kwargs
|
|
790
|
+
Keyword arguments to pass to the to method of each element in the DataList
|
|
791
|
+
|
|
792
|
+
Returns
|
|
793
|
+
-------
|
|
794
|
+
Self
|
|
795
|
+
Self with all elements moved and/or cast
|
|
796
|
+
"""
|
|
797
|
+
return self.apply('to', *args, types=(Tensor, Data), **kwargs)
|
|
798
|
+
|
|
799
|
+
|
|
800
|
+
class BaseDatasetMeta(type):
|
|
801
|
+
"""
|
|
802
|
+
Automatically creates an index for each sample in the dataset after the dataset has been
|
|
803
|
+
initialised.
|
|
804
|
+
"""
|
|
805
|
+
def __call__(cls: type[DatasetT], *args: Any, **kwargs: Any) -> DatasetT:
|
|
806
|
+
"""
|
|
807
|
+
Parameters
|
|
808
|
+
----------
|
|
809
|
+
cls : type[DatasetT]
|
|
810
|
+
Class that inherited BaseDatasetMeta
|
|
811
|
+
*args
|
|
812
|
+
Optional arguments to pass to BaseDataset class
|
|
813
|
+
**kwargs
|
|
814
|
+
Optional keyword arguments to pass to BaseDataset class
|
|
815
|
+
|
|
816
|
+
Returns
|
|
817
|
+
-------
|
|
818
|
+
DatasetT
|
|
819
|
+
Dataset instance from the class that inherited BaseDatasetMeta
|
|
820
|
+
"""
|
|
821
|
+
instance: DatasetT = type.__call__(cls, *args, **kwargs)
|
|
822
|
+
|
|
823
|
+
if hasattr(instance, 'idxs') and instance.idxs.dtype != np.int_:
|
|
824
|
+
raise ValueError(f'idxs attribute already exists and does not have type int '
|
|
825
|
+
f'({instance.idxs.dtype}), idxs attribute must be reserved for sample '
|
|
826
|
+
f'index')
|
|
827
|
+
|
|
828
|
+
if not hasattr(instance, 'high_dim') or instance.high_dim is None:
|
|
829
|
+
raise ValueError(f'{instance.__class__.__name__} has no high_dim attribute which is '
|
|
830
|
+
f'required by BaseDatasetMeta for creating idxs attribute')
|
|
831
|
+
|
|
832
|
+
if len(instance.idxs) == 0:
|
|
833
|
+
instance.idxs = np.arange(len(instance.high_dim))
|
|
834
|
+
elif len(instance.idxs) != len(instance.high_dim):
|
|
835
|
+
log.getLogger(__name__).warning(f'Length of idxs ({len(instance.idxs)}) and length of '
|
|
836
|
+
f'high_dim ({len(instance.high_dim)}) does not match, '
|
|
837
|
+
f'idxs will be sent to a range of high_dim length')
|
|
838
|
+
instance.idxs = np.arange(len(instance.high_dim))
|
|
839
|
+
|
|
840
|
+
for attribute in ('extra', 'low_dim', 'high_dim'):
|
|
841
|
+
if (getattr(instance, attribute) is not None and
|
|
842
|
+
len(getattr(instance, attribute)) != len(instance.idxs)):
|
|
843
|
+
raise ValueError(f'Length of attribute {attribute} '
|
|
844
|
+
f'({len(getattr(instance, attribute))}) and idxs '
|
|
845
|
+
f'({len(instance.idxs)}) does not match')
|
|
846
|
+
return instance
|
|
847
|
+
|
|
848
|
+
|
|
849
|
+
class BaseDataset(Dataset[Any], Generic[DataListT], metaclass=BaseDatasetMeta):
|
|
850
|
+
"""
|
|
851
|
+
Base dataset class for use with BaseNetwork.
|
|
852
|
+
|
|
853
|
+
Attributes
|
|
854
|
+
----------
|
|
855
|
+
extra : list[Any] | ArrayLike | None
|
|
856
|
+
Additional data for each sample in the dataset of length N with shape (N,...) and type Any
|
|
857
|
+
idxs : ndarray
|
|
858
|
+
Index for each sample in the dataset with shape (N) and type int
|
|
859
|
+
low_dim : DataListT | None
|
|
860
|
+
Low dimensional data for each sample in the dataset with shape (N,...)
|
|
861
|
+
high_dim : DataListT | None
|
|
862
|
+
High dimensional data for each sample in the dataset with shape (N,...), this is required
|
|
863
|
+
"""
|
|
864
|
+
def __init__(self) -> None:
|
|
865
|
+
super().__init__()
|
|
866
|
+
self.extra: list[Any] | ArrayLike | None = None
|
|
867
|
+
self.idxs: ndarray = np.array([], dtype=np.int_)
|
|
868
|
+
self.low_dim: DataListT | None = None
|
|
869
|
+
self.high_dim: DataListT | None = None
|
|
870
|
+
|
|
871
|
+
def __len__(self) -> int:
|
|
872
|
+
"""
|
|
873
|
+
Returns the number of samples in the dataset
|
|
874
|
+
|
|
875
|
+
Returns
|
|
876
|
+
-------
|
|
877
|
+
int
|
|
878
|
+
Number of samples in the dataset
|
|
879
|
+
"""
|
|
880
|
+
return len(self.idxs)
|
|
881
|
+
|
|
882
|
+
def __getitem__(self, idx: int) -> tuple[int, DataListT, DataListT, Any]:
|
|
883
|
+
"""
|
|
884
|
+
Parameters
|
|
885
|
+
----------
|
|
886
|
+
idx : int
|
|
887
|
+
Sample index
|
|
888
|
+
|
|
889
|
+
Returns
|
|
890
|
+
-------
|
|
891
|
+
tuple[int, DataListT, DataListT, Any]
|
|
892
|
+
Sample index, low dimensional data, high dimensional data, and extra data
|
|
893
|
+
"""
|
|
894
|
+
return self.idxs[idx], self.get_low_dim(idx), self.get_high_dim(idx), self.get_extra(idx)
|
|
895
|
+
|
|
896
|
+
def get_extra(self, idx: int) -> Any:
|
|
897
|
+
"""
|
|
898
|
+
Gets extra data for the sample of the given index
|
|
899
|
+
|
|
900
|
+
Parameters
|
|
901
|
+
----------
|
|
902
|
+
idx : int
|
|
903
|
+
Sample index
|
|
904
|
+
|
|
905
|
+
Returns
|
|
906
|
+
-------
|
|
907
|
+
Any
|
|
908
|
+
Sample extra data
|
|
909
|
+
"""
|
|
910
|
+
return torch.tensor(()) if self.extra is None else self.extra[idx]
|
|
911
|
+
|
|
912
|
+
def get_high_dim(self, idx: int) -> DataListT:
|
|
913
|
+
"""
|
|
914
|
+
Gets a high dimensional sample of the given index
|
|
915
|
+
|
|
916
|
+
Parameters
|
|
917
|
+
----------
|
|
918
|
+
idx : int
|
|
919
|
+
Sample index
|
|
920
|
+
|
|
921
|
+
Returns
|
|
922
|
+
-------
|
|
923
|
+
DataListT
|
|
924
|
+
High dimensional sample
|
|
925
|
+
"""
|
|
926
|
+
assert self.high_dim is not None
|
|
927
|
+
return cast(DataListT, self.high_dim[idx])
|
|
928
|
+
|
|
929
|
+
def get_low_dim(self, idx: int) -> DataListT:
|
|
930
|
+
"""
|
|
931
|
+
Gets a low dimensional sample of the given index
|
|
932
|
+
|
|
933
|
+
Parameters
|
|
934
|
+
----------
|
|
935
|
+
idx : int
|
|
936
|
+
Sample index
|
|
937
|
+
|
|
938
|
+
Returns
|
|
939
|
+
-------
|
|
940
|
+
DataListT
|
|
941
|
+
Low dimensional sample
|
|
942
|
+
"""
|
|
943
|
+
return cast(DataListT, torch.tensor(()) if self.low_dim is None else self.low_dim[idx])
|
|
944
|
+
|
|
945
|
+
def step(self, epoch: float) -> None:
|
|
946
|
+
"""
|
|
947
|
+
Step the dataset for each training iteration.
|
|
948
|
+
|
|
949
|
+
Parameters
|
|
950
|
+
----------
|
|
951
|
+
epoch : float
|
|
952
|
+
Current epoch number
|
|
953
|
+
"""
|
|
954
|
+
|
|
955
|
+
|
|
956
|
+
@overload
|
|
957
|
+
def data_collation(
|
|
958
|
+
data: list[ArrayCT] | ArrayCT,
|
|
959
|
+
*,
|
|
960
|
+
data_field: bool) -> ArrayCT: ...
|
|
961
|
+
|
|
962
|
+
@overload
|
|
963
|
+
def data_collation(
|
|
964
|
+
data: DataList[DataT],
|
|
965
|
+
*,
|
|
966
|
+
data_field: bool) -> DataList[DataT]: ...
|
|
967
|
+
|
|
968
|
+
@overload
|
|
969
|
+
def data_collation(
|
|
970
|
+
data: list[DataList[DataT]],
|
|
971
|
+
*,
|
|
972
|
+
data_field: Literal[True]) -> DataList[DataT]: ...
|
|
973
|
+
|
|
974
|
+
@overload
|
|
975
|
+
def data_collation(
|
|
976
|
+
data: list[DataList[DataT]],
|
|
977
|
+
*,
|
|
978
|
+
data_field: Literal[False]) -> DataList[ArrayT]: ...
|
|
979
|
+
|
|
980
|
+
@overload
|
|
981
|
+
def data_collation(
|
|
982
|
+
data: list[Data[ArrayCT]],
|
|
983
|
+
*,
|
|
984
|
+
data_field: Literal[True]) -> Data[ArrayCT]: ...
|
|
985
|
+
|
|
986
|
+
@overload
|
|
987
|
+
def data_collation(
|
|
988
|
+
data: list[Data[ArrayCT]],
|
|
989
|
+
*,
|
|
990
|
+
data_field: Literal[False]) -> ArrayCT: ...
|
|
991
|
+
|
|
992
|
+
def data_collation( # type: ignore
|
|
993
|
+
data: list[ArrayCT] | list[Data[ArrayCT]] | list[DataList[DataT]] | ArrayCT |
|
|
994
|
+
DataList[DataT],
|
|
995
|
+
*,
|
|
996
|
+
data_field: bool = True,
|
|
997
|
+
) -> ArrayCT | Data[ArrayCT] | DataList[ArrayT] | DataList[DataT]:
|
|
998
|
+
"""
|
|
999
|
+
Collates a list of ArrayLike, Data, or DataList objects into a single object.
|
|
1000
|
+
|
|
1001
|
+
Parameters
|
|
1002
|
+
----------
|
|
1003
|
+
data : list[ArrayCT] | list[Data[ArrayCT]] | list[DataList[DataT]] | ArrayCT | DataList[DataT]
|
|
1004
|
+
List of ArrayLike, Data, or DataList objects to collate, or if ArrayLike or DataList,
|
|
1005
|
+
will return as is, with shape (N,...)
|
|
1006
|
+
data_field : bool, Optional
|
|
1007
|
+
If Datas should return Data else ArrayLike, default = True
|
|
1008
|
+
|
|
1009
|
+
Returns
|
|
1010
|
+
-------
|
|
1011
|
+
ArrayCT | Data[ArrayCT] | DataList[ArrayT] | DataList[DataT]
|
|
1012
|
+
Collated Data or DataList object, or ArrayLike if data_field is False or input is
|
|
1013
|
+
ArrayLike, with shape (N,...)
|
|
1014
|
+
"""
|
|
1015
|
+
if isinstance(data, (Tensor, ndarray, DataList)):
|
|
1016
|
+
return data
|
|
1017
|
+
if isinstance(data[0], (Tensor, ndarray)):
|
|
1018
|
+
return torch.concat(cast(list[Tensor], data)) if isinstance(data[0], Tensor) else \
|
|
1019
|
+
np.concat(cast(list[ndarray], data))
|
|
1020
|
+
|
|
1021
|
+
# data = cast(list[Data] | list[DataList], data)
|
|
1022
|
+
return Data.collate(cast(list[Data], data), data_field=data_field) \
|
|
1023
|
+
if isinstance(data[0], Data) else \
|
|
1024
|
+
DataList.collate(cast(list[DataList], data), data_field=data_field)
|
|
1025
|
+
|
|
1026
|
+
|
|
1027
|
+
def loader_init(
|
|
1028
|
+
dataset: DatasetT,
|
|
1029
|
+
*,
|
|
1030
|
+
return_idxs: bool = False,
|
|
1031
|
+
batch_size: int = 64,
|
|
1032
|
+
ratios: list[float] | tuple[float, ...] | None = None,
|
|
1033
|
+
idxs: list[ndarray] | tuple[ndarray, ...] | ndarray | None = None,
|
|
1034
|
+
**kwargs: Any) -> (tuple[DataLoader[Subset[DatasetT]], ...] |
|
|
1035
|
+
tuple[tuple[DataLoader[Subset[DatasetT]], ...], tuple[list[int], ...]]):
|
|
1036
|
+
"""
|
|
1037
|
+
Initialises data loaders from a subset of the dataset with the given ratios.
|
|
1038
|
+
|
|
1039
|
+
Parameters
|
|
1040
|
+
----------
|
|
1041
|
+
dataset : DatasetT
|
|
1042
|
+
Dataset to create data loaders from
|
|
1043
|
+
return_idxs : bool, Optional
|
|
1044
|
+
If the indexes for each data loader should be returned, default = False
|
|
1045
|
+
batch_size : int, Optional
|
|
1046
|
+
Batch size when sampling from the data loaders, default = 64
|
|
1047
|
+
ratios : list[float] | tuple[float, ...] | None, Optional
|
|
1048
|
+
Ratios of length M to split up the dataset into subsets, if idxs is provided, dataset will
|
|
1049
|
+
first be split up using idxs and ratios will be used on the remaining samples,
|
|
1050
|
+
default = (0.8,0.2)
|
|
1051
|
+
idxs : list[ndarray] | tuple[ndarray, ...] | ndarray | None, Optional
|
|
1052
|
+
Dataset indexes for creating the subsets with shape (N,S), where N is the number of subsets
|
|
1053
|
+
and S is the number of samples in each subset
|
|
1054
|
+
**kwargs : Any
|
|
1055
|
+
Optional keyword arguments to pass to DataLoader
|
|
1056
|
+
|
|
1057
|
+
Returns
|
|
1058
|
+
-------
|
|
1059
|
+
tuple[DataLoader[Subset[DatasetT]], ...] | tuple[tuple[DataLoader[Subset[DatasetT]], ...], tuple[list[int], ...]]
|
|
1060
|
+
Data loaders for each subset of length N + M and optionally indexes used for each subset
|
|
1061
|
+
"""
|
|
1062
|
+
num: int
|
|
1063
|
+
slice_: float | ndarray
|
|
1064
|
+
loaders: list[DataLoader[Subset[DatasetT]]] = []
|
|
1065
|
+
data_idxs: ndarray = np.arange(len(dataset))
|
|
1066
|
+
loader: DataLoader[Subset[DatasetT]]
|
|
1067
|
+
idxs = [] if idxs is None else list(idxs) \
|
|
1068
|
+
if isinstance(idxs, (tuple, list)) or np.ndim(idxs) > 1 else [idxs]
|
|
1069
|
+
ratios = (0.8, 1) if ratios is None else list(np.cumsum(np.array(ratios) / np.sum(ratios)))
|
|
1070
|
+
np.random.shuffle(data_idxs)
|
|
1071
|
+
|
|
1072
|
+
for slice_ in idxs + list(ratios):
|
|
1073
|
+
if isinstance(slice_, (int, float)):
|
|
1074
|
+
num = max(int(len(data_idxs) * slice_), 1)
|
|
1075
|
+
slice_ = data_idxs[:num]
|
|
1076
|
+
|
|
1077
|
+
if not np.isin(data_idxs, slice_).any():
|
|
1078
|
+
continue
|
|
1079
|
+
|
|
1080
|
+
loaders.append(DataLoader(
|
|
1081
|
+
Subset(dataset, data_idxs[np.isin(data_idxs, slice_)].tolist()),
|
|
1082
|
+
batch_size=batch_size,
|
|
1083
|
+
**{'shuffle': True} | kwargs,
|
|
1084
|
+
))
|
|
1085
|
+
data_idxs = np.delete(data_idxs, np.isin(data_idxs, slice_))
|
|
1086
|
+
|
|
1087
|
+
if return_idxs:
|
|
1088
|
+
return tuple(loaders), tuple(loader.dataset.indices for loader in loaders)
|
|
1089
|
+
return tuple(loaders)
|
|
1090
|
+
|
|
1091
|
+
|
|
1092
|
+
__all__ = ['Data', 'DataList', 'BaseDataset', 'data_collation', 'loader_init']
|