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/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']