netloader 3.11.0__tar.gz → 3.11.2__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 (41) hide show
  1. {netloader-3.11.0 → netloader-3.11.2}/PKG-INFO +7 -9
  2. {netloader-3.11.0 → netloader-3.11.2}/README.md +2 -4
  3. {netloader-3.11.0 → netloader-3.11.2}/netloader/__init__.py +18 -3
  4. {netloader-3.11.0 → netloader-3.11.2}/netloader/architectures/base.py +84 -80
  5. {netloader-3.11.0 → netloader-3.11.2}/netloader/architectures/encoder_decoder.py +11 -11
  6. {netloader-3.11.0 → netloader-3.11.2}/netloader/architectures/flows.py +27 -34
  7. {netloader-3.11.0 → netloader-3.11.2}/netloader/architectures/utils.py +1 -1
  8. {netloader-3.11.0 → netloader-3.11.2}/netloader/data.py +66 -11
  9. {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/base.py +92 -27
  10. {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/misc.py +1 -8
  11. {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/multi_layer.py +14 -15
  12. {netloader-3.11.0 → netloader-3.11.2}/netloader/loss_funcs.py +2 -2
  13. {netloader-3.11.0 → netloader-3.11.2}/netloader/models/convnext.py +8 -8
  14. {netloader-3.11.0 → netloader-3.11.2}/netloader/models/misc.py +56 -24
  15. {netloader-3.11.0 → netloader-3.11.2}/netloader/network.py +207 -178
  16. {netloader-3.11.0 → netloader-3.11.2}/netloader/schedulers.py +12 -1
  17. {netloader-3.11.0 → netloader-3.11.2}/netloader/transforms.py +7 -5
  18. {netloader-3.11.0 → netloader-3.11.2}/netloader/utils/__init__.py +2 -0
  19. {netloader-3.11.0 → netloader-3.11.2}/netloader/utils/types.py +10 -2
  20. {netloader-3.11.0 → netloader-3.11.2}/netloader/utils/utils.py +78 -10
  21. {netloader-3.11.0 → netloader-3.11.2}/netloader.egg-info/PKG-INFO +7 -9
  22. {netloader-3.11.0 → netloader-3.11.2}/netloader.egg-info/requires.txt +3 -3
  23. {netloader-3.11.0 → netloader-3.11.2}/pyproject.toml +4 -4
  24. {netloader-3.11.0 → netloader-3.11.2}/LICENSE.txt +0 -0
  25. {netloader-3.11.0 → netloader-3.11.2}/netloader/architectures/__init__.py +0 -0
  26. {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/__init__.py +0 -0
  27. {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/blocks.py +0 -0
  28. {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/convolutional.py +0 -0
  29. {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/flows.py +0 -0
  30. {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/linear.py +0 -0
  31. {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/pooling.py +0 -0
  32. {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/recurrent.py +0 -0
  33. {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/utils.py +0 -0
  34. {netloader-3.11.0 → netloader-3.11.2}/netloader/models/__init__.py +0 -0
  35. {netloader-3.11.0 → netloader-3.11.2}/netloader/networks/__init__.py +0 -0
  36. {netloader-3.11.0 → netloader-3.11.2}/netloader/utils/configs.py +0 -0
  37. {netloader-3.11.0 → netloader-3.11.2}/netloader/utils/transforms.py +0 -0
  38. {netloader-3.11.0 → netloader-3.11.2}/netloader.egg-info/SOURCES.txt +0 -0
  39. {netloader-3.11.0 → netloader-3.11.2}/netloader.egg-info/dependency_links.txt +0 -0
  40. {netloader-3.11.0 → netloader-3.11.2}/netloader.egg-info/top_level.txt +0 -0
  41. {netloader-3.11.0 → netloader-3.11.2}/setup.cfg +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: netloader
3
- Version: 3.11.0
3
+ Version: 3.11.2
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
@@ -9,14 +9,14 @@ Project-URL: Issues, https://github.com/EthanTreg/PyTorch-Network-Loader/issues
9
9
  Project-URL: ReadTheDocs, https://pytorch-network-loader.readthedocs.io/en/latest/
10
10
  Classifier: Programming Language :: Python :: 3
11
11
  Classifier: Operating System :: OS Independent
12
- Requires-Python: >=3.12
12
+ Requires-Python: >=3.11
13
13
  Description-Content-Type: text/markdown
14
14
  License-File: LICENSE.txt
15
- Requires-Dist: numpy>=2.4.0
16
- Requires-Dist: torch>=2.11.0
15
+ Requires-Dist: numpy>=2.0.0
16
+ Requires-Dist: torch>=2.8.0
17
17
  Requires-Dist: packaging>=24.0
18
18
  Provides-Extra: flows
19
- Requires-Dist: zuko; extra == "flows"
19
+ Requires-Dist: zuko>=1.5.0; extra == "flows"
20
20
  Provides-Extra: wandb
21
21
  Requires-Dist: wandb; extra == "wandb"
22
22
  Dynamic: license-file
@@ -37,16 +37,14 @@ For a real-world example of this package, see [Bayesian-DARKSKIES](https://githu
37
37
 
38
38
  ### Using Within Projects
39
39
 
40
- - pip install `netloader @ git+https://github.com/EthanTreg/PyTorch-Network-Loader@LATEST-VERSION`[^1] to
41
- `requirements.txt`
42
- - Install using `pip install -r requirements.txt`
40
+ - Install using `pip install netloader`[^1]
43
41
  - Example of [InceptionV4](https://arxiv.org/abs/1602.07261) can be downloaded under
44
42
  `./network_configs/inceptionv4.json` along with the composite layers in
45
43
  `./network_configs/composite_layers/`
46
44
 
47
45
  [^1]: To use normalising flows or weights & biases, `netloader` must be pip installed with the optional argument
48
46
  `flows` and/or `wandb`:
49
- `pip install netloader[flows] @ git+https://github.com/EthanTreg/PyTorch-Network-Loader@LATEST-VERSION`
47
+ `pip install netloader[flows,wandb]`
50
48
 
51
49
  ### Locally Running NetLoader
52
50
 
@@ -14,16 +14,14 @@ For a real-world example of this package, see [Bayesian-DARKSKIES](https://githu
14
14
 
15
15
  ### Using Within Projects
16
16
 
17
- - pip install `netloader @ git+https://github.com/EthanTreg/PyTorch-Network-Loader@LATEST-VERSION`[^1] to
18
- `requirements.txt`
19
- - Install using `pip install -r requirements.txt`
17
+ - Install using `pip install netloader`[^1]
20
18
  - Example of [InceptionV4](https://arxiv.org/abs/1602.07261) can be downloaded under
21
19
  `./network_configs/inceptionv4.json` along with the composite layers in
22
20
  `./network_configs/composite_layers/`
23
21
 
24
22
  [^1]: To use normalising flows or weights & biases, `netloader` must be pip installed with the optional argument
25
23
  `flows` and/or `wandb`:
26
- `pip install netloader[flows] @ git+https://github.com/EthanTreg/PyTorch-Network-Loader@LATEST-VERSION`
24
+ `pip install netloader[flows,wandb]`
27
25
 
28
26
  ### Locally Running NetLoader
29
27
 
@@ -6,7 +6,7 @@ import logging
6
6
  import warnings
7
7
 
8
8
 
9
- __version__ = '3.11.0'
9
+ __version__ = '3.11.2'
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(\.|$)')
@@ -43,8 +43,23 @@ try:
43
43
 
44
44
  # Adds PyTorch Network Loader classes to list of safe PyTorch classes when loading saved
45
45
  # architectures
46
- safe_globals(__name__, [models, architectures, loss_funcs, schedulers, transforms])
47
- torch.serialization.add_safe_globals([network.Network, network.CompatibleNetwork])
46
+ safe_globals(
47
+ __name__,
48
+ [models, loss_funcs, schedulers, transforms],
49
+ use_all=True,
50
+ )
51
+ safe_globals(
52
+ __name__,
53
+ [architectures],
54
+ use_all=True,
55
+ legacy_path_replace=('architectures', 'networks'),
56
+ )
57
+ torch.serialization.add_safe_globals([
58
+ network.BaseNetwork,
59
+ network.Network,
60
+ network.CompatibleNetwork,
61
+ ])
62
+ torch.serialization.add_safe_globals([])
48
63
 
49
64
  __all__ = [
50
65
  'utils',
@@ -6,7 +6,8 @@ import logging as log
6
6
  from time import time
7
7
  from warnings import warn
8
8
  from itertools import repeat
9
- from typing import TYPE_CHECKING, Any, Self, Generic, Literal, Callable, cast, overload
9
+ from abc import ABC, abstractmethod
10
+ from typing import TYPE_CHECKING, Any, Self, Sequence, Generic, Literal, Callable, cast, overload
10
11
 
11
12
  import torch
12
13
  import numpy as np
@@ -16,17 +17,12 @@ from torch.utils.data import DataLoader
16
17
  from torch.optim.optimizer import ParamsT
17
18
  from torch._dynamo import OptimizedModule
18
19
 
19
- if TYPE_CHECKING:
20
- from wandb import Run
21
- else:
22
- Run = Any
23
-
24
20
  import netloader
25
21
  from netloader import utils
26
22
  from netloader.transforms import BaseTransform
27
23
  from netloader.architectures.utils import UtilityMixin
28
- from netloader.network import Network, CompatibleNetwork
29
- from netloader.data import Data, DataList, data_collation
24
+ from netloader.data import ApplyFunc, Data, DataList, data_collation
25
+ from netloader.network import BaseNetwork, CompatibleNetwork
30
26
  from netloader.utils.types import (
31
27
  TensorLike,
32
28
  NDArrayLike,
@@ -37,8 +33,13 @@ from netloader.utils.types import (
37
33
  TensorLossCT,
38
34
  )
39
35
 
36
+ if TYPE_CHECKING:
37
+ from wandb import Run
38
+ else:
39
+ Run = Any
40
+
40
41
 
41
- class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
42
+ class BaseArchitecture(UtilityMixin, ABC, Generic[LossCT, TensorLossCT]):
42
43
  # pylint: disable=line-too-long
43
44
  """
44
45
  Base architecture class that other types of architectures build from
@@ -61,15 +62,38 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
61
62
  Architecture optimiser
62
63
  scheduler : LRScheduler
63
64
  Optimiser scheduler
64
- net : Network | CompatibleNetwork
65
+ net : BaseNetwork
65
66
  Neural network
66
67
  """
67
68
  # pylint: enable=line-too-long
69
+ _half: bool
70
+ _train_state: bool = True
71
+ _plot_active: bool = False
72
+ _save_freq: int
73
+ _epoch: int = 0
74
+ _save_path: str = ''
75
+ _verbose: Literal['epoch', 'full', 'plot', 'progress', None]
76
+ _loader_states: tuple[Tensor, Tensor] | None = None
77
+ _loss_weights: dict[str, float | Callable[[], float]]
78
+ _optimiser_kwargs: dict[str, Any]
79
+ _scheduler_kwargs: dict[str, Any]
80
+ _logger: log.Logger = log.getLogger(__name__)
81
+ _device: torch.device = torch.device('cpu')
82
+ description: str
83
+ version: str = netloader.__version__
84
+ losses: tuple[list[LossCT], list[LossCT]]
85
+ transforms: dict[str, Sequence[BaseTransform] | BaseTransform | None]
86
+ idxs: ndarray | None = None
87
+ optimiser: optim.Optimizer
88
+ scheduler: optim.lr_scheduler.LRScheduler
89
+ net: BaseNetwork
90
+ run: Run | None = None
91
+
68
92
  def __init__(
69
93
  self,
70
94
  save_num: int | str,
71
95
  states_dir: str,
72
- net: nn.Module | Network | CompatibleNetwork,
96
+ net: nn.Module | BaseNetwork,
73
97
  *,
74
98
  overwrite: bool = False,
75
99
  mix_precision: bool = False,
@@ -88,7 +112,7 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
88
112
  File number or name to save the architecture
89
113
  states_dir : str
90
114
  Directory to save the architecture
91
- net : Module | Network | CompatibleNetwork
115
+ net : Module | BaseNetwork
92
116
  Network to predict low-dimensional data
93
117
  overwrite : bool, Optional
94
118
  If saving can overwrite an existing save file, if False and file with the same name
@@ -97,7 +121,7 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
97
121
  If mixed precision should be used, default = False
98
122
  save_freq : int, Optional
99
123
  Frequency of epochs to save the architecture, default = 1
100
- learning_rate : float, Optional
124
+ learning_rate : float | tuple[float, ...], Optional
101
125
  Optimiser initial learning rate, default = 1e-3
102
126
  description : str, Optional
103
127
  Description of the architecture
@@ -114,34 +138,21 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
114
138
  scheduler_kwargs : dict[str, Any] | None, Optional
115
139
  Optional keyword arguments to pass to init_scheduler
116
140
  """
117
- self._train_state: bool = True
118
- self._plot_active: bool = False
119
- self._half: bool = mix_precision
120
- self._epoch: int = 0
121
- self._save_freq: int = save_freq
122
- self._save_path: str = ''
123
- self._verbose: Literal['epoch', 'full', 'plot', 'progress', None] = verbose
124
- self._loader_states: tuple[Tensor, Tensor] | None = None
125
- self._loss_weights: dict[str, float | Callable[[], float]] = {}
126
- self._optimiser_kwargs: dict[str, Any] = optimiser_kwargs or {}
127
- self._scheduler_kwargs: dict[str, Any] = scheduler_kwargs or {}
128
- self._logger: log.Logger = log.getLogger(__name__)
129
- self._device: torch.device = torch.device('cpu')
130
- self.description: str = description
131
- self.version: str = netloader.__version__
132
- self.losses: tuple[list[LossCT], list[LossCT]] = ([], [])
133
- self.transforms: dict[str, list[BaseTransform] | BaseTransform | None] = {
141
+ self._half = mix_precision
142
+ self._save_freq = save_freq
143
+ self._verbose = verbose
144
+ self._loss_weights = {}
145
+ self._optimiser_kwargs = optimiser_kwargs or {}
146
+ self._scheduler_kwargs = scheduler_kwargs or {}
147
+ self.description = description
148
+ self.losses = ([], [])
149
+ self.transforms = {
134
150
  'ids': None,
135
151
  'inputs': in_transform,
136
152
  'targets': transform,
137
153
  'preds': transform,
138
154
  }
139
- self.idxs: ndarray | None = None
140
- self.optimiser: optim.Optimizer
141
- self.scheduler: optim.lr_scheduler.LRScheduler
142
- self.net: Network | CompatibleNetwork = net \
143
- if isinstance(net, (Network, CompatibleNetwork)) else CompatibleNetwork(net=net)
144
- self.run: Run | None = None
155
+ self.net = net if isinstance(net, BaseNetwork) else CompatibleNetwork(net=net)
145
156
 
146
157
  if save_num:
147
158
  self._save_path = utils.save_name(save_num, states_dir, self.net.name)
@@ -159,7 +170,7 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
159
170
  DeprecationWarning,
160
171
  stacklevel=2,
161
172
  )
162
- self.init_optimiser = self.set_optimiser
173
+ self.init_optimiser = self.set_optimiser # type: ignore[method-assign]
163
174
 
164
175
  if self._method_override('set_scheduler', UtilityMixin):
165
176
  warn(
@@ -168,7 +179,7 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
168
179
  DeprecationWarning,
169
180
  stacklevel=2,
170
181
  )
171
- self.init_scheduler = self.set_scheduler
182
+ self.init_scheduler = self.set_scheduler # type: ignore[method-assign]
172
183
 
173
184
  self.optimiser = self.init_optimiser(
174
185
  self.get_param_groups(learning_rate),
@@ -177,10 +188,11 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
177
188
  self.scheduler = self.init_scheduler(self.optimiser, **self._scheduler_kwargs)
178
189
 
179
190
  if isinstance(self.scheduler, optim.lr_scheduler.ReduceLROnPlateau):
180
- self._scheduler_kwargs = ({
181
- 'factor': 0.5,
182
- 'min_lr': (learning_rate if isinstance(learning_rate, float) else learning_rate[0]) * 1e-3,
183
- } | self._scheduler_kwargs)
191
+ self._scheduler_kwargs = {
192
+ 'factor': 0.5,
193
+ 'min_lr': (learning_rate if isinstance(learning_rate, float) else learning_rate[0])
194
+ * 1e-3
195
+ } | self._scheduler_kwargs
184
196
  self.scheduler.load_state_dict(self._scheduler_kwargs)
185
197
 
186
198
  # Adds all architecture classes to list of safe PyTorch classes when loading saved
@@ -296,7 +308,7 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
296
308
  DeprecationWarning,
297
309
  stacklevel=2,
298
310
  )
299
- self.init_optimiser = self.set_optimiser
311
+ self.init_optimiser = self.set_optimiser # type: ignore[method-assign]
300
312
 
301
313
  if self._method_override('set_scheduler', UtilityMixin):
302
314
  warn(
@@ -305,7 +317,7 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
305
317
  DeprecationWarning,
306
318
  stacklevel=2,
307
319
  )
308
- self.init_scheduler = self.set_scheduler
320
+ self.init_scheduler = self.set_scheduler # type: ignore[method-assign]
309
321
 
310
322
  if isinstance(state['optimiser'], dict):
311
323
  self.optimiser = self.init_optimiser(
@@ -494,7 +506,7 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
494
506
  DeprecationWarning,
495
507
  stacklevel=2,
496
508
  )
497
- loss = self._loss_tensor(in_data, target)
509
+ loss = self._loss_tensor(in_data, target) # type: ignore[call-arg]
498
510
 
499
511
  if isinstance(loss, dict) and 'total' not in loss:
500
512
  loss['total'] = self._loss_total(loss)
@@ -523,6 +535,7 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
523
535
  """
524
536
  raise DeprecationWarning
525
537
 
538
+ @abstractmethod
526
539
  def _loss_tensor(
527
540
  self,
528
541
  in_data: TensorListLike,
@@ -545,7 +558,6 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
545
558
  TensorLossCT
546
559
  Loss of shape (1) and type float or dictionary of losses of shape (1) and type float
547
560
  """
548
- raise NotImplementedError
549
561
 
550
562
  def _loss_total(self, losses: dict[str, Tensor]) -> Tensor:
551
563
  """
@@ -666,7 +678,7 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
666
678
  DeprecationWarning,
667
679
  stacklevel=2,
668
680
  )
669
- batch_loss = self._loss(in_data, target)
681
+ batch_loss = self._loss(in_data, target) # type: ignore[call-arg] # pylint: disable=no-value-for-parameter
670
682
 
671
683
  if isinstance(batch_loss, dict) and loss:
672
684
  for key, value in batch_loss.items():
@@ -680,27 +692,14 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
680
692
  loss = batch_loss
681
693
 
682
694
  if self._train_state:
683
- try:
684
- self._step(
685
- True,
686
- self._epoch + (i + 1) / len(loader),
687
- batch_loss['total'] if isinstance(batch_loss, dict) else batch_loss,
688
- dataset=loader.dataset.dataset
689
- if hasattr(loader, 'dataset') and hasattr(loader.dataset, 'dataset')
690
- else None,
691
- )
692
- except TypeError:
693
- warn(
694
- '_step without dataset parameter is deprecated, please update'
695
- 'the method to include the dataset parameter',
696
- DeprecationWarning,
697
- stacklevel=2,
698
- )
699
- self._step(
700
- True,
701
- self._epoch + (i + 1) / len(loader),
702
- batch_loss['total'] if isinstance(batch_loss, dict) else batch_loss,
703
- )
695
+ self._step(
696
+ True,
697
+ self._epoch + (i + 1) / len(loader),
698
+ batch_loss['total'] if isinstance(batch_loss, dict) else batch_loss,
699
+ dataset=loader.dataset.dataset
700
+ if hasattr(loader, 'dataset') and hasattr(loader.dataset, 'dataset')
701
+ else None,
702
+ )
704
703
 
705
704
  self._batch_print(i, time() - t_initial, loader, batch_loss)
706
705
 
@@ -746,8 +745,8 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
746
745
  metrics : float | None, Optional
747
746
  Loss metric to update ReduceLROnPlateau
748
747
  """
749
- learning_rate: list[float]
750
- new_learning_rate: list[float]
748
+ learning_rate: list[float | Tensor]
749
+ new_learning_rate: list[float | Tensor]
751
750
 
752
751
  if batch_step is None:
753
752
  warn(
@@ -812,9 +811,9 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
812
811
  Optional keyword arguments to pass to torch.compile
813
812
  """
814
813
  if level == 'net':
815
- self.net = cast(Network | CompatibleNetwork, torch.compile(self.net, **kwargs))
814
+ self.net = cast(BaseNetwork, torch.compile(self.net, **kwargs))
816
815
  elif level == 'loss':
817
- self._loss_tensor = torch.compile(self._loss_tensor, **kwargs) # type: ignore[method-assign] # pylint: disable=attribute-defined-outside-init
816
+ self._loss_tensor = torch.compile(self._loss_tensor, **kwargs) # type: ignore[method-assign] # pylint: disable=attribute-defined-outside-init, method-hidden
818
817
  else:
819
818
  raise ValueError(f'Invalid compile level ({level}), must be "net" or "loss"')
820
819
 
@@ -1023,11 +1022,13 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
1023
1022
  high_dim: list[Tensor] | list[Data[Tensor]] | list[DataList[Tensor | Data[Tensor]]]
1024
1023
  data: list[list[NDArrayLike | None]] = []
1025
1024
  data_: dict[str, NDArrayLike] = {}
1026
- transform: list[BaseTransform] | BaseTransform | None
1027
- transforms: dict[str, list[BaseTransform] | BaseTransform | None] = {
1028
- key: transform for key, transform in self.transforms.items()
1029
- if inputs or key != 'inputs'
1030
- }
1025
+ transform: (Sequence[BaseTransform[ndarray, ndarray]] |
1026
+ BaseTransform[ndarray, ndarray] | None)
1027
+ transforms: dict[
1028
+ str,
1029
+ Sequence[BaseTransform[ndarray, ndarray]] | BaseTransform[ndarray, ndarray] | None
1030
+ ] = {key: transform for key, transform in self.transforms.items()
1031
+ if inputs or key != 'inputs'}
1031
1032
  datum: NDArrayLike
1032
1033
  target: TensorLike
1033
1034
  in_data: TensorLike
@@ -1075,7 +1076,10 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
1075
1076
  datum = data_collation(list(datum_), data_field=True)
1076
1077
 
1077
1078
  if isinstance(datum, DataList) and transform:
1078
- data_[key] = datum.apply(transform, back=True).numpy()
1079
+ data_[key] = datum.apply(cast(
1080
+ ApplyFunc[ndarray | Data[ndarray], ..., ndarray | Data[ndarray]],
1081
+ transform,
1082
+ ), back=True).numpy()
1079
1083
  data_[key] = DataList(
1080
1084
  [trans(val, back=True) for val, trans in zip(
1081
1085
  datum,
@@ -1198,7 +1202,7 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
1198
1202
  cuda_state: Tensor | None
1199
1203
  np.random.set_state(states.pop('numpy', np.random.get_state()))
1200
1204
  self._loader_states = states.pop('loaders', None)
1201
- torch.set_rng_state(states.pop('pytorch', torch.get_rng_state()))
1205
+ torch.set_rng_state(states.pop('pytorch', torch.get_rng_state()).cpu())
1202
1206
 
1203
1207
  if torch.cuda.is_available() and (cuda_state := states.pop('cuda_rng', None)):
1204
1208
  torch.cuda.set_rng_state_all(cuda_state)
@@ -1222,7 +1226,7 @@ class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
1222
1226
  self._device, *_ = torch._C._nn._parse_to(*args, **kwargs) # pylint: disable=protected-access
1223
1227
 
1224
1228
  if isinstance(self.net, OptimizedModule):
1225
- self.net._orig_mod.to(*args, **kwargs)
1229
+ self.net._orig_mod.to(*args, **kwargs) # pylint: disable=protected-access
1226
1230
  return self
1227
1231
 
1228
1232
  def train(self, train: bool) -> None:
@@ -11,8 +11,8 @@ from torch import Tensor, nn
11
11
 
12
12
  from netloader.data import DataList
13
13
  from netloader.utils import label_change
14
+ from netloader.network import BaseNetwork
14
15
  from netloader.transforms import BaseTransform
15
- from netloader.network import Network, CompatibleNetwork
16
16
  from netloader.architectures.base import BaseArchitecture
17
17
  from netloader.loss_funcs import BaseLoss, MSELoss, CrossEntropyLoss
18
18
  from netloader.utils.types import TensorListLike, NDArrayListLike, TensorT
@@ -42,14 +42,14 @@ class BaseEncoder(BaseArchitecture):
42
42
  Architecture optimiser
43
43
  scheduler : LRScheduler
44
44
  Optimiser scheduler
45
- net : Network | CompatibleNetwork
45
+ net : BaseNetwork
46
46
  Neural network
47
47
  """
48
48
  def __init__(
49
49
  self,
50
50
  save_num: int | str,
51
51
  states_dir: str,
52
- net: nn.Module | Network | CompatibleNetwork,
52
+ net: nn.Module | BaseNetwork,
53
53
  *,
54
54
  overwrite: bool = False,
55
55
  mix_precision: bool = False,
@@ -69,7 +69,7 @@ class BaseEncoder(BaseArchitecture):
69
69
  File number or name to save the architecture
70
70
  states_dir : str
71
71
  Directory to save the architecture
72
- net : Module | Network | CompatibleNetwork
72
+ net : Module | BaseNetwork
73
73
  Network to predict low-dimensional data
74
74
  overwrite : bool, Optional
75
75
  If saving can overwrite an existing save file, if True and file with the same name
@@ -223,7 +223,7 @@ class Autoencoder(BaseArchitecture):
223
223
  Architecture optimiser
224
224
  scheduler : LRScheduler
225
225
  Optimiser scheduler
226
- net : Network | CompatibleNetwork
226
+ net : BaseNetwork
227
227
  Neural network
228
228
  reconstruct_func : BaseLoss
229
229
  Loss function for the reconstruction loss
@@ -234,7 +234,7 @@ class Autoencoder(BaseArchitecture):
234
234
  self,
235
235
  save_num: int | str,
236
236
  states_dir: str,
237
- net: nn.Module | Network | CompatibleNetwork,
237
+ net: nn.Module | BaseNetwork,
238
238
  *,
239
239
  overwrite: bool = False,
240
240
  mix_precision: bool = False,
@@ -252,7 +252,7 @@ class Autoencoder(BaseArchitecture):
252
252
  File number or name to save the architecture
253
253
  states_dir : str
254
254
  Directory to save the architecture
255
- net : Module | Network | CompatibleNetwork
255
+ net : Module | BaseNetwork
256
256
  Network to predict low-dimensional data
257
257
  overwrite : bool, Optional
258
258
  If saving can overwrite an existing save file, if True and file with the same name
@@ -441,14 +441,14 @@ class Decoder(BaseArchitecture):
441
441
  Architecture optimiser
442
442
  scheduler : LRScheduler
443
443
  Optimiser scheduler
444
- net : Network | CompatibleNetwork
444
+ net : BaseNetwork
445
445
  Neural network
446
446
  """
447
447
  def __init__(
448
448
  self,
449
449
  save_num: int | str,
450
450
  states_dir: str,
451
- net: nn.Module | Network | CompatibleNetwork,
451
+ net: nn.Module | BaseNetwork,
452
452
  *,
453
453
  overwrite: bool = False,
454
454
  mix_precision: bool = False,
@@ -466,7 +466,7 @@ class Decoder(BaseArchitecture):
466
466
  File number or name to save the architecture
467
467
  states_dir : str
468
468
  Directory to save the architecture
469
- net : Module | Network | CompatibleNetwork
469
+ net : Module | BaseNetwork
470
470
  Network to predict low-dimensional data
471
471
  overwrite : bool, Optional
472
472
  If saving can overwrite an existing save file, if True and file with the same name
@@ -601,7 +601,7 @@ class Encoder(BaseEncoder):
601
601
  Architecture optimiser
602
602
  scheduler : LRScheduler
603
603
  Optimiser scheduler
604
- net : Network | CompatibleNetwork
604
+ net : BaseNetwork
605
605
  Neural network
606
606
  """
607
607
  def batch_predict(self, data: TensorListLike, **_: Any) -> tuple[NDArrayListLike | None, ...]:
@@ -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
17
18
  from netloader.transforms import BaseTransform
18
- from netloader.network import Network, CompatibleNetwork
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,7 +29,7 @@ class NormFlow(BaseArchitecture):
29
29
 
30
30
  Attributes
31
31
  ----------
32
- net : Network | CompatibleNetwork
32
+ net : BaseNetwork
33
33
  Neural spline flow
34
34
  description : str
35
35
  Description of the architecture
@@ -146,14 +146,20 @@ class NormFlowEncoder(BaseEncoder):
146
146
  Architecture optimiser
147
147
  scheduler : LRScheduler
148
148
  Optimiser scheduler
149
- net : Network | CompatibleNetwork
149
+ net : BaseNetwork
150
150
  Neural network
151
151
  """
152
+ _train_flow: bool
153
+ _train_encoder: bool
154
+ _checkpoint: int | None
155
+ _epochs: tuple[int, int]
156
+ net: BaseNetwork[nn.ModuleList]
157
+
152
158
  def __init__(
153
159
  self,
154
160
  save_num: int | str,
155
161
  states_dir: str,
156
- net: nn.Module | Network | CompatibleNetwork,
162
+ net: nn.Module | BaseNetwork[nn.ModuleList],
157
163
  *,
158
164
  overwrite: bool = False,
159
165
  mix_precision: bool = False,
@@ -175,7 +181,7 @@ class NormFlowEncoder(BaseEncoder):
175
181
  File number or name to save the flow
176
182
  states_dir : str
177
183
  Directory to save the architecture and flow
178
- net : Module | Network | CompatibleNetwork
184
+ net : Module | BaseNetwork[ModuleList]
179
185
  Normalising flow to predict low-dimensional data distribution
180
186
  overwrite : bool, Optional
181
187
  If saving can overwrite an existing save file, if True and file with the same name
@@ -226,11 +232,8 @@ class NormFlowEncoder(BaseEncoder):
226
232
  optimiser_kwargs=optimiser_kwargs,
227
233
  scheduler_kwargs=scheduler_kwargs,
228
234
  )
229
- self._train_flow: bool
230
- self._train_encoder: bool
231
- self._checkpoint: int | None = net_checkpoint
232
- self._epochs: tuple[int, int] = train_epochs
233
-
235
+ self._checkpoint = net_checkpoint
236
+ self._epochs = train_epochs
234
237
  self._train_flow = not self._epochs[0]
235
238
  self._train_encoder = bool(self._epochs[-1])
236
239
  self._loss_weights = {'flow': 1, 'encoder': 1}
@@ -240,6 +243,7 @@ class NormFlowEncoder(BaseEncoder):
240
243
  'max': transform,
241
244
  'meds': transform,
242
245
  }
246
+ assert isinstance(self.net, nn.ModuleList)
243
247
 
244
248
  if not self._train_encoder:
245
249
  self.net.layers[:-1].requires_grad_(False)
@@ -360,8 +364,8 @@ class NormFlowEncoder(BaseEncoder):
360
364
  Returns
361
365
  -------
362
366
  tuple[NDArrayListLike | None, ndarray | None]
363
- Architecture output with shape (N,...) and type float and samples of shape (N,S) and type
364
- float from each probability distribution
367
+ Architecture output with shape (N,...) and type float and samples of shape (N,S) and
368
+ type float from each probability distribution
365
369
  """
366
370
  samples: ndarray | None = None
367
371
  output: NormalizingFlow | TensorListLike | None = self.net(data)
@@ -411,6 +415,17 @@ class NormFlowEncoder(BaseEncoder):
411
415
  'train_epochs': self._epochs,
412
416
  }
413
417
 
418
+ def get_param_groups(self, learning_rate: float | tuple[float, ...] | None) -> ParamsT:
419
+ learning_rate = learning_rate or (0,) * 2
420
+
421
+ if isinstance(learning_rate, float):
422
+ learning_rate = (learning_rate,) * 2
423
+
424
+ return [
425
+ {'params': self.net.layers[:-1].parameters(), 'lr': learning_rate[0]},
426
+ {'params': self.net.layers[-1:].parameters(), 'lr': learning_rate[1]},
427
+ ]
428
+
414
429
  def predict(
415
430
  self,
416
431
  loader: DataLoader[Any],
@@ -475,27 +490,5 @@ class NormFlowEncoder(BaseEncoder):
475
490
  self._save_predictions(path, data)
476
491
  return data
477
492
 
478
- def get_param_groups(self, learning_rate: tuple[float, float] | None) -> ParamsT:
479
- """
480
- Sets the optimiser for the architecture, by default AdamW.
481
-
482
- Parameters
483
- ----------
484
- learning_rate : tuple[float, float] | None, Optional
485
- Learning rate for the encoder and normalising flow, if None, learning rate is set to 0
486
- **kwargs
487
- Optional keyword arguments to pass to the optimiser
488
-
489
- Returns
490
- -------
491
- Optimizer
492
- Architecture optimiser
493
- """
494
- learning_rate = learning_rate or (0,) * 2
495
- return [
496
- {'params': self.net.layers[:-1].parameters(), 'lr': learning_rate[0]},
497
- {'params': self.net.layers[-1:].parameters(), 'lr': learning_rate[1]},
498
- ]
499
-
500
493
 
501
494
  __all__ = ['NormFlow', 'NormFlowEncoder']
@@ -98,7 +98,7 @@ class UtilityMixin:
98
98
  bool
99
99
  True if the method is overridden, False otherwise
100
100
  """
101
- method: Callable = getattr(self, name, None)
101
+ method: Callable | None = getattr(self, name, None)
102
102
 
103
103
  if method is None:
104
104
  logger.warning(f'Method {name} not found in {type(self).__name__}')