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.
- {netloader-3.11.0 → netloader-3.11.2}/PKG-INFO +7 -9
- {netloader-3.11.0 → netloader-3.11.2}/README.md +2 -4
- {netloader-3.11.0 → netloader-3.11.2}/netloader/__init__.py +18 -3
- {netloader-3.11.0 → netloader-3.11.2}/netloader/architectures/base.py +84 -80
- {netloader-3.11.0 → netloader-3.11.2}/netloader/architectures/encoder_decoder.py +11 -11
- {netloader-3.11.0 → netloader-3.11.2}/netloader/architectures/flows.py +27 -34
- {netloader-3.11.0 → netloader-3.11.2}/netloader/architectures/utils.py +1 -1
- {netloader-3.11.0 → netloader-3.11.2}/netloader/data.py +66 -11
- {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/base.py +92 -27
- {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/misc.py +1 -8
- {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/multi_layer.py +14 -15
- {netloader-3.11.0 → netloader-3.11.2}/netloader/loss_funcs.py +2 -2
- {netloader-3.11.0 → netloader-3.11.2}/netloader/models/convnext.py +8 -8
- {netloader-3.11.0 → netloader-3.11.2}/netloader/models/misc.py +56 -24
- {netloader-3.11.0 → netloader-3.11.2}/netloader/network.py +207 -178
- {netloader-3.11.0 → netloader-3.11.2}/netloader/schedulers.py +12 -1
- {netloader-3.11.0 → netloader-3.11.2}/netloader/transforms.py +7 -5
- {netloader-3.11.0 → netloader-3.11.2}/netloader/utils/__init__.py +2 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader/utils/types.py +10 -2
- {netloader-3.11.0 → netloader-3.11.2}/netloader/utils/utils.py +78 -10
- {netloader-3.11.0 → netloader-3.11.2}/netloader.egg-info/PKG-INFO +7 -9
- {netloader-3.11.0 → netloader-3.11.2}/netloader.egg-info/requires.txt +3 -3
- {netloader-3.11.0 → netloader-3.11.2}/pyproject.toml +4 -4
- {netloader-3.11.0 → netloader-3.11.2}/LICENSE.txt +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader/architectures/__init__.py +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/__init__.py +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/blocks.py +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/convolutional.py +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/flows.py +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/linear.py +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/pooling.py +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/recurrent.py +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader/layers/utils.py +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader/models/__init__.py +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader/networks/__init__.py +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader/utils/configs.py +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader/utils/transforms.py +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader.egg-info/SOURCES.txt +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader.egg-info/dependency_links.txt +0 -0
- {netloader-3.11.0 → netloader-3.11.2}/netloader.egg-info/top_level.txt +0 -0
- {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.
|
|
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
|
+
Requires-Python: >=3.11
|
|
13
13
|
Description-Content-Type: text/markdown
|
|
14
14
|
License-File: LICENSE.txt
|
|
15
|
-
Requires-Dist: numpy>=2.
|
|
16
|
-
Requires-Dist: torch>=2.
|
|
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
|
|
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]
|
|
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
|
|
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]
|
|
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.
|
|
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(
|
|
47
|
-
|
|
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
|
|
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.
|
|
29
|
-
from netloader.
|
|
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 :
|
|
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 |
|
|
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 |
|
|
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.
|
|
118
|
-
self.
|
|
119
|
-
self.
|
|
120
|
-
self.
|
|
121
|
-
self.
|
|
122
|
-
self.
|
|
123
|
-
self.
|
|
124
|
-
self.
|
|
125
|
-
self.
|
|
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.
|
|
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
|
-
|
|
182
|
-
|
|
183
|
-
|
|
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
|
-
|
|
684
|
-
|
|
685
|
-
|
|
686
|
-
|
|
687
|
-
|
|
688
|
-
|
|
689
|
-
|
|
690
|
-
|
|
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(
|
|
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:
|
|
1027
|
-
|
|
1028
|
-
|
|
1029
|
-
|
|
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(
|
|
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 :
|
|
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 |
|
|
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 |
|
|
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 :
|
|
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 |
|
|
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 |
|
|
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 :
|
|
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 |
|
|
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 |
|
|
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 :
|
|
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 :
|
|
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 :
|
|
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 |
|
|
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 |
|
|
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.
|
|
230
|
-
self.
|
|
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
|
|
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__}')
|