netloader 3.11.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- netloader/__init__.py +61 -0
- netloader/architectures/__init__.py +48 -0
- netloader/architectures/base.py +1350 -0
- netloader/architectures/encoder_decoder.py +630 -0
- netloader/architectures/flows.py +501 -0
- netloader/architectures/utils.py +179 -0
- netloader/data.py +1092 -0
- netloader/layers/__init__.py +88 -0
- netloader/layers/base.py +367 -0
- netloader/layers/blocks.py +145 -0
- netloader/layers/convolutional.py +1008 -0
- netloader/layers/flows.py +123 -0
- netloader/layers/linear.py +445 -0
- netloader/layers/misc.py +879 -0
- netloader/layers/multi_layer.py +470 -0
- netloader/layers/pooling.py +298 -0
- netloader/layers/recurrent.py +163 -0
- netloader/layers/utils.py +168 -0
- netloader/loss_funcs.py +201 -0
- netloader/models/__init__.py +22 -0
- netloader/models/convnext.py +582 -0
- netloader/models/misc.py +209 -0
- netloader/network.py +895 -0
- netloader/networks/__init__.py +48 -0
- netloader/schedulers.py +409 -0
- netloader/transforms.py +847 -0
- netloader/utils/__init__.py +46 -0
- netloader/utils/configs.py +241 -0
- netloader/utils/transforms.py +22 -0
- netloader/utils/types.py +130 -0
- netloader/utils/utils.py +677 -0
- netloader-3.11.0.dist-info/METADATA +59 -0
- netloader-3.11.0.dist-info/RECORD +36 -0
- netloader-3.11.0.dist-info/WHEEL +5 -0
- netloader-3.11.0.dist-info/licenses/LICENSE.txt +21 -0
- netloader-3.11.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,1350 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Base architecture class to base other architectures off
|
|
3
|
+
"""
|
|
4
|
+
import os
|
|
5
|
+
import logging as log
|
|
6
|
+
from time import time
|
|
7
|
+
from warnings import warn
|
|
8
|
+
from itertools import repeat
|
|
9
|
+
from typing import TYPE_CHECKING, Any, Self, Generic, Literal, Callable, cast, overload
|
|
10
|
+
|
|
11
|
+
import torch
|
|
12
|
+
import numpy as np
|
|
13
|
+
from numpy import ndarray
|
|
14
|
+
from torch import nn, optim, Tensor
|
|
15
|
+
from torch.utils.data import DataLoader
|
|
16
|
+
from torch.optim.optimizer import ParamsT
|
|
17
|
+
from torch._dynamo import OptimizedModule
|
|
18
|
+
|
|
19
|
+
if TYPE_CHECKING:
|
|
20
|
+
from wandb import Run
|
|
21
|
+
else:
|
|
22
|
+
Run = Any
|
|
23
|
+
|
|
24
|
+
import netloader
|
|
25
|
+
from netloader import utils
|
|
26
|
+
from netloader.transforms import BaseTransform
|
|
27
|
+
from netloader.architectures.utils import UtilityMixin
|
|
28
|
+
from netloader.network import Network, CompatibleNetwork
|
|
29
|
+
from netloader.data import Data, DataList, data_collation
|
|
30
|
+
from netloader.utils.types import (
|
|
31
|
+
TensorLike,
|
|
32
|
+
NDArrayLike,
|
|
33
|
+
TensorListLike,
|
|
34
|
+
NDArrayListLike,
|
|
35
|
+
DatasetT,
|
|
36
|
+
LossCT,
|
|
37
|
+
TensorLossCT,
|
|
38
|
+
)
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class BaseArchitecture(UtilityMixin, Generic[LossCT, TensorLossCT]):
|
|
42
|
+
# pylint: disable=line-too-long
|
|
43
|
+
"""
|
|
44
|
+
Base architecture class that other types of architectures build from
|
|
45
|
+
|
|
46
|
+
Attributes
|
|
47
|
+
----------
|
|
48
|
+
description : str
|
|
49
|
+
Description of the architecture
|
|
50
|
+
version : str
|
|
51
|
+
Version of the architecture when it was created or re-saved
|
|
52
|
+
losses : tuple[list[LossCT], list[LossCT]]
|
|
53
|
+
Architecture training and validation losses as a float or dictionary of losses for each loss
|
|
54
|
+
function
|
|
55
|
+
transforms : dict[str, list[BaseTransform] | BaseTransform | None]
|
|
56
|
+
Keys for the output data from predict and corresponding transforms
|
|
57
|
+
idxs: ndarray | None
|
|
58
|
+
Training data indices with shape (N) and type int, where N is the number of elements in the
|
|
59
|
+
training dataset
|
|
60
|
+
optimiser : Optimizer
|
|
61
|
+
Architecture optimiser
|
|
62
|
+
scheduler : LRScheduler
|
|
63
|
+
Optimiser scheduler
|
|
64
|
+
net : Network | CompatibleNetwork
|
|
65
|
+
Neural network
|
|
66
|
+
"""
|
|
67
|
+
# pylint: enable=line-too-long
|
|
68
|
+
def __init__(
|
|
69
|
+
self,
|
|
70
|
+
save_num: int | str,
|
|
71
|
+
states_dir: str,
|
|
72
|
+
net: nn.Module | Network | CompatibleNetwork,
|
|
73
|
+
*,
|
|
74
|
+
overwrite: bool = False,
|
|
75
|
+
mix_precision: bool = False,
|
|
76
|
+
save_freq: int = 1,
|
|
77
|
+
learning_rate: float | tuple[float, ...] = 1e-3,
|
|
78
|
+
description: str = '',
|
|
79
|
+
verbose: Literal['epoch', 'full', 'plot', 'progress', None] = 'epoch',
|
|
80
|
+
transform: list[BaseTransform] | BaseTransform | None = None,
|
|
81
|
+
in_transform: list[BaseTransform] | BaseTransform | None = None,
|
|
82
|
+
optimiser_kwargs: dict[str, Any] | None = None,
|
|
83
|
+
scheduler_kwargs: dict[str, Any] | None = None) -> None:
|
|
84
|
+
"""
|
|
85
|
+
Parameters
|
|
86
|
+
----------
|
|
87
|
+
save_num : int | str
|
|
88
|
+
File number or name to save the architecture
|
|
89
|
+
states_dir : str
|
|
90
|
+
Directory to save the architecture
|
|
91
|
+
net : Module | Network | CompatibleNetwork
|
|
92
|
+
Network to predict low-dimensional data
|
|
93
|
+
overwrite : bool, Optional
|
|
94
|
+
If saving can overwrite an existing save file, if False and file with the same name
|
|
95
|
+
exists, an error will be raised, default = False
|
|
96
|
+
mix_precision : bool, Optional
|
|
97
|
+
If mixed precision should be used, default = False
|
|
98
|
+
save_freq : int, Optional
|
|
99
|
+
Frequency of epochs to save the architecture, default = 1
|
|
100
|
+
learning_rate : float, Optional
|
|
101
|
+
Optimiser initial learning rate, default = 1e-3
|
|
102
|
+
description : str, Optional
|
|
103
|
+
Description of the architecture
|
|
104
|
+
verbose : {'epoch', 'full', 'plot', 'progress', None}
|
|
105
|
+
If details about each epoch should be printed ('epoch'), details about epoch and epoch
|
|
106
|
+
progress (full), details about epoch and an ASCII plot of the loss progress ('plot'),
|
|
107
|
+
just total progress ('progress'), or nothing (None)
|
|
108
|
+
transform : list[BaseTransform] | BaseTransform | None, Optional
|
|
109
|
+
Transformation(s) of the network's output(s)
|
|
110
|
+
in_transform : list[BaseTransform] | BaseTransform | None, Optional
|
|
111
|
+
Transformation(s) for the network's input(s)
|
|
112
|
+
optimiser_kwargs : dict[str, Any] | None, Optional
|
|
113
|
+
Optional keyword arguments to pass to init_optimiser
|
|
114
|
+
scheduler_kwargs : dict[str, Any] | None, Optional
|
|
115
|
+
Optional keyword arguments to pass to init_scheduler
|
|
116
|
+
"""
|
|
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] = {
|
|
134
|
+
'ids': None,
|
|
135
|
+
'inputs': in_transform,
|
|
136
|
+
'targets': transform,
|
|
137
|
+
'preds': transform,
|
|
138
|
+
}
|
|
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
|
|
145
|
+
|
|
146
|
+
if save_num:
|
|
147
|
+
self._save_path = utils.save_name(save_num, states_dir, self.net.name)
|
|
148
|
+
|
|
149
|
+
if os.path.exists(self._save_path) and overwrite:
|
|
150
|
+
self._logger.warning(f'{self._save_path} already exists and will be overwritten if '
|
|
151
|
+
f'training continues')
|
|
152
|
+
elif os.path.exists(self._save_path):
|
|
153
|
+
raise FileExistsError(f'{self._save_path} already exists and overwrite is False')
|
|
154
|
+
|
|
155
|
+
if self._method_override('set_optimiser', UtilityMixin):
|
|
156
|
+
warn(
|
|
157
|
+
'BaseArchitecture.set_optimiser method is deprecated, please override '
|
|
158
|
+
'init_optimiser instead',
|
|
159
|
+
DeprecationWarning,
|
|
160
|
+
stacklevel=2,
|
|
161
|
+
)
|
|
162
|
+
self.init_optimiser = self.set_optimiser
|
|
163
|
+
|
|
164
|
+
if self._method_override('set_scheduler', UtilityMixin):
|
|
165
|
+
warn(
|
|
166
|
+
'BaseArchitecture.set_scheduler method is deprecated, please override '
|
|
167
|
+
'init_scheduler instead',
|
|
168
|
+
DeprecationWarning,
|
|
169
|
+
stacklevel=2,
|
|
170
|
+
)
|
|
171
|
+
self.init_scheduler = self.set_scheduler
|
|
172
|
+
|
|
173
|
+
self.optimiser = self.init_optimiser(
|
|
174
|
+
self.get_param_groups(learning_rate),
|
|
175
|
+
**self._optimiser_kwargs,
|
|
176
|
+
)
|
|
177
|
+
self.scheduler = self.init_scheduler(self.optimiser, **self._scheduler_kwargs)
|
|
178
|
+
|
|
179
|
+
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)
|
|
184
|
+
self.scheduler.load_state_dict(self._scheduler_kwargs)
|
|
185
|
+
|
|
186
|
+
# Adds all architecture classes to list of safe PyTorch classes when loading saved
|
|
187
|
+
# architectures
|
|
188
|
+
torch.serialization.add_safe_globals([self.__class__])
|
|
189
|
+
|
|
190
|
+
def __repr__(self) -> str:
|
|
191
|
+
"""
|
|
192
|
+
Returns a string representation of the architecture.
|
|
193
|
+
|
|
194
|
+
Returns
|
|
195
|
+
-------
|
|
196
|
+
str
|
|
197
|
+
String representation of the architecture
|
|
198
|
+
"""
|
|
199
|
+
return (f'Architecture: {self.__class__.__name__}\n'
|
|
200
|
+
f'Name: {os.path.basename(self._save_path).rsplit(".", 1)[0]}\n'
|
|
201
|
+
f'Description: {self.description}\n'
|
|
202
|
+
f'Version: {self.version}\n'
|
|
203
|
+
f'Network: {self.net.name}\n'
|
|
204
|
+
f'Epoch: {self._epoch}\n'
|
|
205
|
+
f'Optimiser: {self.optimiser.__class__.__name__}\n'
|
|
206
|
+
f'Scheduler: {self.scheduler.__class__.__name__ if self.scheduler else None}\n' +
|
|
207
|
+
(f'Loss Weights: {self.get_loss_weights()}\n' if self._loss_weights else '') +
|
|
208
|
+
f'Args: ({self.extra_repr()})')
|
|
209
|
+
|
|
210
|
+
def __getstate__(self) -> dict[str, Any]:
|
|
211
|
+
"""
|
|
212
|
+
Returns a dictionary containing the state of the architecture for pickling.
|
|
213
|
+
|
|
214
|
+
Returns
|
|
215
|
+
-------
|
|
216
|
+
dict[str, Any]
|
|
217
|
+
Dictionary containing the state of the architecture
|
|
218
|
+
"""
|
|
219
|
+
return {
|
|
220
|
+
'half': self._half,
|
|
221
|
+
'epoch': self._epoch,
|
|
222
|
+
'save_freq': self._save_freq,
|
|
223
|
+
'verbose': self._verbose,
|
|
224
|
+
'save_path': self._save_path,
|
|
225
|
+
'description': self.description,
|
|
226
|
+
'version': netloader.__version__,
|
|
227
|
+
'losses': self.losses,
|
|
228
|
+
'loss_weights': self._loss_weights,
|
|
229
|
+
'optimiser_kwargs': self._optimiser_kwargs,
|
|
230
|
+
'scheduler_kwargs': self._scheduler_kwargs,
|
|
231
|
+
'transforms': self.transforms,
|
|
232
|
+
'rng_states': self.get_rng_states(),
|
|
233
|
+
'idxs': None if self.idxs is None else self.idxs.tolist(),
|
|
234
|
+
'optimiser': self.optimiser.state_dict(),
|
|
235
|
+
'scheduler': self.scheduler.state_dict(),
|
|
236
|
+
'net': self.net._orig_mod if isinstance(self.net, OptimizedModule) else self.net,
|
|
237
|
+
}
|
|
238
|
+
|
|
239
|
+
def __setstate__(self, state: dict[str, Any]) -> None:
|
|
240
|
+
"""
|
|
241
|
+
Sets the state of the architecture for pickling.
|
|
242
|
+
|
|
243
|
+
Parameters
|
|
244
|
+
----------
|
|
245
|
+
state : dict[str, Any]
|
|
246
|
+
Dictionary containing the state of the architecture
|
|
247
|
+
"""
|
|
248
|
+
self._train_state = True
|
|
249
|
+
self._plot_active = False
|
|
250
|
+
|
|
251
|
+
for key, value in list(state.items()):
|
|
252
|
+
if key[0] == '_':
|
|
253
|
+
state[key.replace('_', '', 1)] = value
|
|
254
|
+
|
|
255
|
+
self._half = state['half']
|
|
256
|
+
self._epoch = state['epoch']
|
|
257
|
+
self._save_freq = state.get('save_freq', 1)
|
|
258
|
+
self._verbose = state['verbose']
|
|
259
|
+
self._save_path = state['save_path']
|
|
260
|
+
self._loader_states = None
|
|
261
|
+
self._loss_weights = state.get('loss_weights', {})
|
|
262
|
+
self._optimiser_kwargs = state.get('optimiser_kwargs', {})
|
|
263
|
+
self._scheduler_kwargs = state.get('scheduler_kwargs', {})
|
|
264
|
+
self._logger = log.getLogger(__name__)
|
|
265
|
+
self._device = torch.device('cpu')
|
|
266
|
+
self.version = state.get('version', '<3.7.1')
|
|
267
|
+
self.description = state['description']
|
|
268
|
+
self.losses = state['losses']
|
|
269
|
+
self.transforms = state['transforms'] if 'transforms' in state else state['header']
|
|
270
|
+
self.idxs = state['idxs'] if state['idxs'] is None else np.array(state['idxs'])
|
|
271
|
+
self.net = state['net']
|
|
272
|
+
self.run = state.get('run', None)
|
|
273
|
+
|
|
274
|
+
if not utils.compare_versions(self.version, netloader.__version__):
|
|
275
|
+
warn(
|
|
276
|
+
f'Architecture version ({self.version}) is older than the current '
|
|
277
|
+
f'NetLoader version ({netloader.__version__}), please resave the architecture '
|
|
278
|
+
f'using BaseArchitecture.save()',
|
|
279
|
+
DeprecationWarning,
|
|
280
|
+
stacklevel=2,
|
|
281
|
+
)
|
|
282
|
+
|
|
283
|
+
if 'header' in state:
|
|
284
|
+
warn(
|
|
285
|
+
'header attribute of BaseArchitecture is deprecated, please resave the '
|
|
286
|
+
'architecture with the new attribute name using BaseArchitecture.save()',
|
|
287
|
+
DeprecationWarning,
|
|
288
|
+
stacklevel=2,
|
|
289
|
+
)
|
|
290
|
+
self.transforms['inputs'] = state['in_transform']
|
|
291
|
+
|
|
292
|
+
if self._method_override('set_optimiser', UtilityMixin):
|
|
293
|
+
warn(
|
|
294
|
+
'BaseArchitecture.set_optimiser method is deprecated, please override '
|
|
295
|
+
'init_optimiser instead',
|
|
296
|
+
DeprecationWarning,
|
|
297
|
+
stacklevel=2,
|
|
298
|
+
)
|
|
299
|
+
self.init_optimiser = self.set_optimiser
|
|
300
|
+
|
|
301
|
+
if self._method_override('set_scheduler', UtilityMixin):
|
|
302
|
+
warn(
|
|
303
|
+
'BaseArchitecture.set_scheduler method is deprecated, please override '
|
|
304
|
+
'init_scheduler instead',
|
|
305
|
+
DeprecationWarning,
|
|
306
|
+
stacklevel=2,
|
|
307
|
+
)
|
|
308
|
+
self.init_scheduler = self.set_scheduler
|
|
309
|
+
|
|
310
|
+
if isinstance(state['optimiser'], dict):
|
|
311
|
+
self.optimiser = self.init_optimiser(
|
|
312
|
+
self.get_param_groups(None),
|
|
313
|
+
**state.get('optimiser_kwargs', {}),
|
|
314
|
+
)
|
|
315
|
+
self.scheduler = self.init_scheduler(
|
|
316
|
+
self.optimiser,
|
|
317
|
+
**state.get('scheduler_kwargs', {}),
|
|
318
|
+
)
|
|
319
|
+
self.optimiser.load_state_dict(state['optimiser'])
|
|
320
|
+
self.scheduler.load_state_dict(state['scheduler'])
|
|
321
|
+
else:
|
|
322
|
+
warn(
|
|
323
|
+
'Optimiser & scheduler is saved in old non-weights safe format and is '
|
|
324
|
+
'deprecated, please resave the architecture in the new format using '
|
|
325
|
+
'BaseArchitecture.save()',
|
|
326
|
+
DeprecationWarning,
|
|
327
|
+
stacklevel=2,
|
|
328
|
+
)
|
|
329
|
+
self.optimiser = state['optimiser']
|
|
330
|
+
self.scheduler = state['scheduler']
|
|
331
|
+
|
|
332
|
+
if 'rng_states' in state:
|
|
333
|
+
self.set_rng_states(state['rng_states'])
|
|
334
|
+
|
|
335
|
+
def __setattr__(self, key: str, value: Any) -> None:
|
|
336
|
+
"""
|
|
337
|
+
Sets attributes of the architecture, with deprecation warning for save_path.
|
|
338
|
+
|
|
339
|
+
Parameters
|
|
340
|
+
----------
|
|
341
|
+
key : str
|
|
342
|
+
Attribute name
|
|
343
|
+
value : Any
|
|
344
|
+
Attribute value
|
|
345
|
+
"""
|
|
346
|
+
if key == 'save_path':
|
|
347
|
+
warn(
|
|
348
|
+
'save_path attribute is deprecated, please use set_save_path method '
|
|
349
|
+
'instead',
|
|
350
|
+
DeprecationWarning,
|
|
351
|
+
stacklevel=2,
|
|
352
|
+
)
|
|
353
|
+
self.set_save_path(str(value), overwrite=True)
|
|
354
|
+
else:
|
|
355
|
+
super().__setattr__(key, value)
|
|
356
|
+
|
|
357
|
+
def __getattr__(self, key: str) -> Any:
|
|
358
|
+
"""
|
|
359
|
+
Gets attributes of the architecture, with deprecation warning for save_path.
|
|
360
|
+
|
|
361
|
+
Parameters
|
|
362
|
+
----------
|
|
363
|
+
key : str
|
|
364
|
+
Attribute name
|
|
365
|
+
|
|
366
|
+
Returns
|
|
367
|
+
-------
|
|
368
|
+
Any
|
|
369
|
+
Attribute value
|
|
370
|
+
"""
|
|
371
|
+
if key == 'save_path':
|
|
372
|
+
warn(
|
|
373
|
+
'save_path attribute is deprecated, please use get_save_path method '
|
|
374
|
+
'instead',
|
|
375
|
+
DeprecationWarning,
|
|
376
|
+
stacklevel=2,
|
|
377
|
+
)
|
|
378
|
+
return self._save_path
|
|
379
|
+
raise AttributeError(f"'{self.__class__.__name__}' object has no attribute '{key}'")
|
|
380
|
+
|
|
381
|
+
def _batch_print(
|
|
382
|
+
self,
|
|
383
|
+
i: int,
|
|
384
|
+
batch_time: float,
|
|
385
|
+
loader: DataLoader[Any],
|
|
386
|
+
loss: LossCT) -> None:
|
|
387
|
+
"""
|
|
388
|
+
Print function during each batch of training.
|
|
389
|
+
|
|
390
|
+
Parameters
|
|
391
|
+
----------
|
|
392
|
+
i : int
|
|
393
|
+
Batch number
|
|
394
|
+
batch_time : float
|
|
395
|
+
Time taken for the batch
|
|
396
|
+
loader: DataLoader[Any]
|
|
397
|
+
Data loader that the batch came from
|
|
398
|
+
loss : LossCT
|
|
399
|
+
Loss for the batch or a dictionary of losses
|
|
400
|
+
"""
|
|
401
|
+
if self._verbose == 'full':
|
|
402
|
+
loss = (cast(dict, loss)['total'] if isinstance(loss, dict) else loss) / (i + 1)
|
|
403
|
+
utils.progress_bar(
|
|
404
|
+
i,
|
|
405
|
+
len(loader),
|
|
406
|
+
text=f"Average loss: {loss:.2e}\tTime: {batch_time:.1f}",
|
|
407
|
+
)
|
|
408
|
+
|
|
409
|
+
def _epoch_print(
|
|
410
|
+
self,
|
|
411
|
+
i: int,
|
|
412
|
+
epochs: int,
|
|
413
|
+
epoch_time: float) -> None:
|
|
414
|
+
"""
|
|
415
|
+
Print function at the end of each epoch.
|
|
416
|
+
|
|
417
|
+
Parameters
|
|
418
|
+
----------
|
|
419
|
+
i : int
|
|
420
|
+
Current epoch number
|
|
421
|
+
epochs : int
|
|
422
|
+
Total number of epochs
|
|
423
|
+
epoch_time : float
|
|
424
|
+
Time taken for the epoch
|
|
425
|
+
"""
|
|
426
|
+
text: str
|
|
427
|
+
losses: tuple[list[float], list[float]] = (
|
|
428
|
+
[cast(dict, loss)['total']
|
|
429
|
+
if isinstance(loss, dict) else loss for loss in self.losses[0]],
|
|
430
|
+
[cast(dict, loss)['total']
|
|
431
|
+
if isinstance(loss, dict) else loss for loss in self.losses[1]],
|
|
432
|
+
)
|
|
433
|
+
loss: LossCT
|
|
434
|
+
|
|
435
|
+
text = f'Epoch [{self._epoch}/{epochs}]\t' \
|
|
436
|
+
f'Training loss: {losses[0][-1]:.3e}\t' \
|
|
437
|
+
f'Validation loss: {losses[1][-1]:.3e}\t' \
|
|
438
|
+
f'Time: {epoch_time:.1f}'
|
|
439
|
+
|
|
440
|
+
if (self._verbose in {'full', 'epoch'} or
|
|
441
|
+
(len(losses[0]) == 1 and self._verbose == 'plot')):
|
|
442
|
+
print(text)
|
|
443
|
+
elif self._verbose == 'progress':
|
|
444
|
+
utils.progress_bar(i, epochs, text=text)
|
|
445
|
+
elif self._verbose == 'plot':
|
|
446
|
+
utils.ascii_plot(
|
|
447
|
+
losses[0],
|
|
448
|
+
clear=self._plot_active,
|
|
449
|
+
text=text,
|
|
450
|
+
data2=losses[1],
|
|
451
|
+
)
|
|
452
|
+
self._plot_active = True
|
|
453
|
+
|
|
454
|
+
def _loss(self, in_data: TensorListLike, target: TensorListLike, extra: Any) -> LossCT:
|
|
455
|
+
"""
|
|
456
|
+
Returns the loss as a float & updates network weights if training.
|
|
457
|
+
|
|
458
|
+
Parameters
|
|
459
|
+
----------
|
|
460
|
+
in_data : TensorListLike
|
|
461
|
+
Input data of shape (N,...) and type float, where N is the number of elements
|
|
462
|
+
target : TensorListLike
|
|
463
|
+
Target data of shape (N,...) and type float
|
|
464
|
+
extra : Any
|
|
465
|
+
Extra data from the data loader to pass to the loss function
|
|
466
|
+
|
|
467
|
+
Returns
|
|
468
|
+
-------
|
|
469
|
+
LossCT
|
|
470
|
+
Loss or dictionary of losses which can be summed to get the total loss
|
|
471
|
+
"""
|
|
472
|
+
key: str
|
|
473
|
+
value: Tensor
|
|
474
|
+
loss: TensorLossCT
|
|
475
|
+
|
|
476
|
+
with torch.autocast(
|
|
477
|
+
enabled=self._half,
|
|
478
|
+
dtype=torch.bfloat16 if self._device == torch.device('cpu') else torch.float16,
|
|
479
|
+
device_type=self._device.type):
|
|
480
|
+
try:
|
|
481
|
+
loss = self._loss_func(in_data, target)
|
|
482
|
+
warn(
|
|
483
|
+
'_loss_func is deprecated, please use _loss_tensor instead',
|
|
484
|
+
DeprecationWarning,
|
|
485
|
+
stacklevel=2,
|
|
486
|
+
)
|
|
487
|
+
except DeprecationWarning:
|
|
488
|
+
try:
|
|
489
|
+
loss = self._loss_tensor(in_data, target, extra)
|
|
490
|
+
except TypeError:
|
|
491
|
+
warn(
|
|
492
|
+
'_loss_tensor without extra parameter is deprecated, please update '
|
|
493
|
+
'the method to include the extra parameter',
|
|
494
|
+
DeprecationWarning,
|
|
495
|
+
stacklevel=2,
|
|
496
|
+
)
|
|
497
|
+
loss = self._loss_tensor(in_data, target)
|
|
498
|
+
|
|
499
|
+
if isinstance(loss, dict) and 'total' not in loss:
|
|
500
|
+
loss['total'] = self._loss_total(loss)
|
|
501
|
+
|
|
502
|
+
self._update(loss['total'] if isinstance(loss, dict) else loss)
|
|
503
|
+
|
|
504
|
+
if isinstance(loss, dict):
|
|
505
|
+
return {key: value.item() for key, value in loss.items()} # type: ignore[return-value]
|
|
506
|
+
return loss.item() # type: ignore[return-value]
|
|
507
|
+
|
|
508
|
+
def _loss_func(self, in_data: TensorListLike, target: TensorListLike) -> TensorLossCT:
|
|
509
|
+
"""
|
|
510
|
+
Empty method for child classes to base their loss functions on.
|
|
511
|
+
|
|
512
|
+
Parameters
|
|
513
|
+
----------
|
|
514
|
+
in_data : TensorListLike
|
|
515
|
+
Input data of shape (N,...) and type float, where N is the number of elements
|
|
516
|
+
target : TensorListLike
|
|
517
|
+
Target data of shape (N,...) and type float
|
|
518
|
+
|
|
519
|
+
Returns
|
|
520
|
+
-------
|
|
521
|
+
TensorLossCT
|
|
522
|
+
Loss of shape (1) and type float or dictionary of losses of shape (1) and type float
|
|
523
|
+
"""
|
|
524
|
+
raise DeprecationWarning
|
|
525
|
+
|
|
526
|
+
def _loss_tensor(
|
|
527
|
+
self,
|
|
528
|
+
in_data: TensorListLike,
|
|
529
|
+
target: TensorListLike,
|
|
530
|
+
extra: Any) -> TensorLossCT:
|
|
531
|
+
"""
|
|
532
|
+
Empty method for child classes to base their loss functions on.
|
|
533
|
+
|
|
534
|
+
Parameters
|
|
535
|
+
----------
|
|
536
|
+
in_data : TensorListLike
|
|
537
|
+
Input data of shape (N,...) and type float, where N is the number of elements
|
|
538
|
+
target : TensorListLike
|
|
539
|
+
Target data of shape (N,...) and type float
|
|
540
|
+
extra : Any
|
|
541
|
+
Extra data from the data loader
|
|
542
|
+
|
|
543
|
+
Returns
|
|
544
|
+
-------
|
|
545
|
+
TensorLossCT
|
|
546
|
+
Loss of shape (1) and type float or dictionary of losses of shape (1) and type float
|
|
547
|
+
"""
|
|
548
|
+
raise NotImplementedError
|
|
549
|
+
|
|
550
|
+
def _loss_total(self, losses: dict[str, Tensor]) -> Tensor:
|
|
551
|
+
"""
|
|
552
|
+
Weighted sum of the loss function terms.
|
|
553
|
+
|
|
554
|
+
Parameters
|
|
555
|
+
----------
|
|
556
|
+
losses : dict[str, Tensor]
|
|
557
|
+
Dictionary of losses of shape (1) and type float
|
|
558
|
+
|
|
559
|
+
Returns
|
|
560
|
+
-------
|
|
561
|
+
Tensor
|
|
562
|
+
Total loss of shape (1) and type float
|
|
563
|
+
"""
|
|
564
|
+
key: str
|
|
565
|
+
loss_weight: float | Callable[[], float]
|
|
566
|
+
val: Tensor
|
|
567
|
+
total_loss: Tensor = torch.tensor(0., device=self._device)
|
|
568
|
+
|
|
569
|
+
for key, val in losses.items():
|
|
570
|
+
loss_weight = self._loss_weights.get(key, 1.0)
|
|
571
|
+
total_loss += val * (loss_weight() if callable(loss_weight) else loss_weight)
|
|
572
|
+
return total_loss
|
|
573
|
+
|
|
574
|
+
def _predict_print(self, i: int, predict_time: float, loader: DataLoader[Any]) -> None:
|
|
575
|
+
"""
|
|
576
|
+
Print function during each batch of prediction.
|
|
577
|
+
|
|
578
|
+
Parameters
|
|
579
|
+
----------
|
|
580
|
+
i : int
|
|
581
|
+
Batch number
|
|
582
|
+
predict_time : float
|
|
583
|
+
Time taken for predicting
|
|
584
|
+
loader: DataLoader[Any]
|
|
585
|
+
Data loader that is being used for predicting
|
|
586
|
+
"""
|
|
587
|
+
if self._verbose == 'full':
|
|
588
|
+
utils.progress_bar(i, len(loader))
|
|
589
|
+
|
|
590
|
+
if i == len(loader) - 1 and self._verbose is not None:
|
|
591
|
+
print(f'Prediction time: {predict_time:.3e} s')
|
|
592
|
+
|
|
593
|
+
def _step(
|
|
594
|
+
self,
|
|
595
|
+
batch_step: bool,
|
|
596
|
+
epoch: int | float,
|
|
597
|
+
metric: float,
|
|
598
|
+
dataset: DatasetT | None = None) -> None:
|
|
599
|
+
"""
|
|
600
|
+
Step method that is called each iteration and each epoch.
|
|
601
|
+
|
|
602
|
+
Parameters
|
|
603
|
+
----------
|
|
604
|
+
batch_step : bool
|
|
605
|
+
If the step is being called during a batch or at the end of an epoch
|
|
606
|
+
epoch : int | float
|
|
607
|
+
Current epoch number, if batch_step is False else current epoch + batch progress
|
|
608
|
+
metric : float
|
|
609
|
+
Loss metric for the iteration or epoch
|
|
610
|
+
dataset : DatasetT, Optional
|
|
611
|
+
Dataset used for training
|
|
612
|
+
"""
|
|
613
|
+
loss: float | object
|
|
614
|
+
self._update_scheduler(batch_step=batch_step, epoch=epoch, metrics=metric)
|
|
615
|
+
|
|
616
|
+
for loss in self._loss_weights.values():
|
|
617
|
+
if hasattr(loss, 'step'):
|
|
618
|
+
loss.step(epoch)
|
|
619
|
+
|
|
620
|
+
if dataset and hasattr(dataset, 'step'):
|
|
621
|
+
dataset.step(epoch)
|
|
622
|
+
|
|
623
|
+
def _train_val(self, loader: DataLoader[Any]) -> LossCT:
|
|
624
|
+
"""
|
|
625
|
+
Trains the network for one epoch.
|
|
626
|
+
|
|
627
|
+
Parameters
|
|
628
|
+
----------
|
|
629
|
+
loader : DataLoader
|
|
630
|
+
PyTorch DataLoader that contains data to train
|
|
631
|
+
|
|
632
|
+
Returns
|
|
633
|
+
-------
|
|
634
|
+
LossCT
|
|
635
|
+
Average loss value or dictionary of average loss values for the epoch
|
|
636
|
+
"""
|
|
637
|
+
i: int
|
|
638
|
+
value: float
|
|
639
|
+
t_initial: float
|
|
640
|
+
key: str
|
|
641
|
+
loss: LossCT | None = None
|
|
642
|
+
low_dim: list[Tensor] | list[Data[Tensor]] | list[DataList[Tensor | Data[Tensor]]]
|
|
643
|
+
high_dim: list[Tensor] | list[Data[Tensor]] | list[DataList[Tensor | Data[Tensor]]]
|
|
644
|
+
extra: list[Any]
|
|
645
|
+
batch_loss: LossCT
|
|
646
|
+
target: TensorListLike
|
|
647
|
+
in_data: TensorListLike
|
|
648
|
+
|
|
649
|
+
with torch.set_grad_enabled(self._train_state):
|
|
650
|
+
for i, (_, low_dim, high_dim, *extra) in enumerate(loader):
|
|
651
|
+
t_initial = time()
|
|
652
|
+
in_data, target = cast(
|
|
653
|
+
tuple[TensorListLike, TensorListLike],
|
|
654
|
+
self._data_loader_translation(
|
|
655
|
+
data_collation(low_dim, data_field=False).to(self._device),
|
|
656
|
+
data_collation(high_dim, data_field=False).to(self._device),
|
|
657
|
+
),
|
|
658
|
+
)
|
|
659
|
+
|
|
660
|
+
try:
|
|
661
|
+
batch_loss = self._loss(in_data, target, extra[0] if extra else None)
|
|
662
|
+
except TypeError:
|
|
663
|
+
warn(
|
|
664
|
+
'_loss without extra parameter is deprecated, please update the '
|
|
665
|
+
'method to include the extra parameter',
|
|
666
|
+
DeprecationWarning,
|
|
667
|
+
stacklevel=2,
|
|
668
|
+
)
|
|
669
|
+
batch_loss = self._loss(in_data, target)
|
|
670
|
+
|
|
671
|
+
if isinstance(batch_loss, dict) and loss:
|
|
672
|
+
for key, value in batch_loss.items():
|
|
673
|
+
if key in loss:
|
|
674
|
+
loss[key] += value
|
|
675
|
+
else:
|
|
676
|
+
loss[key] = value
|
|
677
|
+
elif isinstance(batch_loss, float) and loss:
|
|
678
|
+
loss += batch_loss
|
|
679
|
+
else:
|
|
680
|
+
loss = batch_loss
|
|
681
|
+
|
|
682
|
+
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
|
+
)
|
|
704
|
+
|
|
705
|
+
self._batch_print(i, time() - t_initial, loader, batch_loss)
|
|
706
|
+
|
|
707
|
+
assert loss
|
|
708
|
+
|
|
709
|
+
if isinstance(loss, dict):
|
|
710
|
+
return {key: value / len(loader) for key, value in loss.items()}
|
|
711
|
+
return loss / len(loader)
|
|
712
|
+
|
|
713
|
+
def _update(self, loss: Tensor) -> None:
|
|
714
|
+
"""
|
|
715
|
+
Updates the network using backpropagation.
|
|
716
|
+
|
|
717
|
+
Parameters
|
|
718
|
+
----------
|
|
719
|
+
loss : Tensor
|
|
720
|
+
Loss to perform backpropagation from
|
|
721
|
+
"""
|
|
722
|
+
if self._train_state:
|
|
723
|
+
self.optimiser.zero_grad()
|
|
724
|
+
loss.backward()
|
|
725
|
+
self.optimiser.step()
|
|
726
|
+
|
|
727
|
+
def _update_epoch(self) -> None:
|
|
728
|
+
"""
|
|
729
|
+
Updates architecture epoch.
|
|
730
|
+
"""
|
|
731
|
+
self._epoch += 1
|
|
732
|
+
|
|
733
|
+
def _update_scheduler(
|
|
734
|
+
self,
|
|
735
|
+
batch_step: bool | None = None,
|
|
736
|
+
*,
|
|
737
|
+
metrics: float | None = None,
|
|
738
|
+
**_: Any) -> None:
|
|
739
|
+
"""
|
|
740
|
+
Updates the scheduler for the architecture.
|
|
741
|
+
|
|
742
|
+
Parameters
|
|
743
|
+
----------
|
|
744
|
+
batch_step : bool
|
|
745
|
+
If the step is being called during a batch or at the end of an epoch
|
|
746
|
+
metrics : float | None, Optional
|
|
747
|
+
Loss metric to update ReduceLROnPlateau
|
|
748
|
+
"""
|
|
749
|
+
learning_rate: list[float]
|
|
750
|
+
new_learning_rate: list[float]
|
|
751
|
+
|
|
752
|
+
if batch_step is None:
|
|
753
|
+
warn(
|
|
754
|
+
'batch_step is now a required argument for _update_scheduler, please update'
|
|
755
|
+
'the method call',
|
|
756
|
+
DeprecationWarning,
|
|
757
|
+
stacklevel=2,
|
|
758
|
+
)
|
|
759
|
+
batch_step = metrics is not None
|
|
760
|
+
|
|
761
|
+
if batch_step and not isinstance(self.scheduler, optim.lr_scheduler.ReduceLROnPlateau):
|
|
762
|
+
self.scheduler.step()
|
|
763
|
+
return
|
|
764
|
+
if batch_step or not isinstance(self.scheduler, optim.lr_scheduler.ReduceLROnPlateau):
|
|
765
|
+
return
|
|
766
|
+
|
|
767
|
+
if metrics is None:
|
|
768
|
+
warn(
|
|
769
|
+
'metrics is a required argument for _update_scheduler when batch_step is '
|
|
770
|
+
'False and scheduler is ReduceLROnPlateau, please update the method call',
|
|
771
|
+
DeprecationWarning,
|
|
772
|
+
stacklevel=2,
|
|
773
|
+
)
|
|
774
|
+
return
|
|
775
|
+
|
|
776
|
+
try:
|
|
777
|
+
learning_rate = self.scheduler.get_last_lr()
|
|
778
|
+
self.scheduler.step(metrics)
|
|
779
|
+
new_learning_rate = self.scheduler.get_last_lr()
|
|
780
|
+
|
|
781
|
+
if learning_rate[-1] != new_learning_rate[-1]:
|
|
782
|
+
print(f'Learning rate update: {new_learning_rate[-1]:.3e}')
|
|
783
|
+
except AttributeError:
|
|
784
|
+
self.scheduler.step(metrics)
|
|
785
|
+
|
|
786
|
+
def batch_predict(self, data: TensorListLike, **_: Any) -> tuple[NDArrayListLike | None, ...]:
|
|
787
|
+
"""
|
|
788
|
+
Generates predictions for the given data batch.
|
|
789
|
+
|
|
790
|
+
Parameters
|
|
791
|
+
----------
|
|
792
|
+
data : TensorListLike
|
|
793
|
+
Data of shape (N,...) and type float to generate predictions for, where N is the batch
|
|
794
|
+
size
|
|
795
|
+
|
|
796
|
+
Returns
|
|
797
|
+
-------
|
|
798
|
+
tuple[NDArrayListLike | None, ...]
|
|
799
|
+
Predictions of shape (N,...) and type float for the given data
|
|
800
|
+
"""
|
|
801
|
+
return (self.net(data).detach().cpu().numpy(),)
|
|
802
|
+
|
|
803
|
+
def compile(self, level: Literal['net', 'loss'] = 'net', **kwargs: Any) -> None:
|
|
804
|
+
"""
|
|
805
|
+
Compiles the network using torch.compile for faster training and prediction.
|
|
806
|
+
|
|
807
|
+
Parameters
|
|
808
|
+
----------
|
|
809
|
+
level : {'net', 'loss'}
|
|
810
|
+
If 'net', compiles the network, if 'loss', compiles the loss function
|
|
811
|
+
**kwargs
|
|
812
|
+
Optional keyword arguments to pass to torch.compile
|
|
813
|
+
"""
|
|
814
|
+
if level == 'net':
|
|
815
|
+
self.net = cast(Network | CompatibleNetwork, torch.compile(self.net, **kwargs))
|
|
816
|
+
elif level == 'loss':
|
|
817
|
+
self._loss_tensor = torch.compile(self._loss_tensor, **kwargs) # type: ignore[method-assign] # pylint: disable=attribute-defined-outside-init
|
|
818
|
+
else:
|
|
819
|
+
raise ValueError(f'Invalid compile level ({level}), must be "net" or "loss"')
|
|
820
|
+
|
|
821
|
+
def extra_repr(self) -> str:
|
|
822
|
+
"""
|
|
823
|
+
Additional representation of the architecture.
|
|
824
|
+
|
|
825
|
+
Returns
|
|
826
|
+
-------
|
|
827
|
+
str
|
|
828
|
+
Architecture specific representation
|
|
829
|
+
"""
|
|
830
|
+
return ''
|
|
831
|
+
|
|
832
|
+
def get_device(self) -> torch.device:
|
|
833
|
+
"""
|
|
834
|
+
Gets the device of the architecture.
|
|
835
|
+
|
|
836
|
+
Returns
|
|
837
|
+
-------
|
|
838
|
+
torch.device
|
|
839
|
+
Device of the architecture
|
|
840
|
+
"""
|
|
841
|
+
return self._device
|
|
842
|
+
|
|
843
|
+
def get_epochs(self) -> int:
|
|
844
|
+
"""
|
|
845
|
+
Returns the number of epochs the architecture has been trained for.
|
|
846
|
+
|
|
847
|
+
Returns
|
|
848
|
+
-------
|
|
849
|
+
int
|
|
850
|
+
Number of epochs
|
|
851
|
+
"""
|
|
852
|
+
return self._epoch
|
|
853
|
+
|
|
854
|
+
def get_hyperparams(self) -> dict[str, Any]:
|
|
855
|
+
"""
|
|
856
|
+
Returns the hyperparameters of the architecture.
|
|
857
|
+
|
|
858
|
+
Returns
|
|
859
|
+
-------
|
|
860
|
+
dict[str, Any]
|
|
861
|
+
Hyperparameters of the architecture
|
|
862
|
+
"""
|
|
863
|
+
return {
|
|
864
|
+
'mix_precision': self._half,
|
|
865
|
+
'save_freq': self._save_freq,
|
|
866
|
+
'description': self.description,
|
|
867
|
+
'net_name': self.net.name,
|
|
868
|
+
'architecture_name': self.__class__.__name__,
|
|
869
|
+
'version': self.version,
|
|
870
|
+
'verbose': self._verbose,
|
|
871
|
+
'optimiser': self.optimiser.__class__.__name__,
|
|
872
|
+
'scheduler': self.scheduler.__class__.__name__,
|
|
873
|
+
'loss_weights': self._loss_weights,
|
|
874
|
+
'optimiser_kwargs': self._optimiser_kwargs,
|
|
875
|
+
'scheduler_kwargs': self._scheduler_kwargs,
|
|
876
|
+
}
|
|
877
|
+
|
|
878
|
+
def get_loader_states(self) -> tuple[Tensor, Tensor] | None:
|
|
879
|
+
"""
|
|
880
|
+
Gets the current data loader states.
|
|
881
|
+
|
|
882
|
+
Returns
|
|
883
|
+
-------
|
|
884
|
+
tuple[Tensor, Tensor] | None
|
|
885
|
+
Train and validation data loader states
|
|
886
|
+
"""
|
|
887
|
+
return self._loader_states
|
|
888
|
+
|
|
889
|
+
def get_losses(
|
|
890
|
+
self
|
|
891
|
+
) -> tuple[list[float], list[float]] | tuple[dict[str, list[float]], dict[str, list[float]]]:
|
|
892
|
+
"""
|
|
893
|
+
Returns the training and validation losses as dictionaries of losses if loss is a
|
|
894
|
+
dictionary, else returns losses.
|
|
895
|
+
|
|
896
|
+
Returns
|
|
897
|
+
-------
|
|
898
|
+
tuple[list[float], list[float]] | tuple[dict[str, list[float]], dict[str, list[float]]]
|
|
899
|
+
Training and validation losses as dictionaries of lists if loss is a dictionary, else as
|
|
900
|
+
lists
|
|
901
|
+
"""
|
|
902
|
+
if not isinstance(self.losses[0][0], dict):
|
|
903
|
+
return self.losses
|
|
904
|
+
return (
|
|
905
|
+
utils.list_dict_convert(self.losses[0], concat=True),
|
|
906
|
+
utils.list_dict_convert(self.losses[1], concat=True),
|
|
907
|
+
)
|
|
908
|
+
|
|
909
|
+
@overload
|
|
910
|
+
def get_loss_weights(self, name: str) -> float:
|
|
911
|
+
...
|
|
912
|
+
|
|
913
|
+
@overload
|
|
914
|
+
def get_loss_weights(self) -> dict[str, float]:
|
|
915
|
+
...
|
|
916
|
+
|
|
917
|
+
def get_loss_weights(self, name: str = '') -> float | dict[str, float]:
|
|
918
|
+
"""
|
|
919
|
+
Gets the weights for a loss function term or all loss function term weights.
|
|
920
|
+
|
|
921
|
+
Parameters
|
|
922
|
+
----------
|
|
923
|
+
name : str, Optional
|
|
924
|
+
Name of the loss term to get the weight for, if empty string, returns all loss term
|
|
925
|
+
weights
|
|
926
|
+
|
|
927
|
+
Returns
|
|
928
|
+
-------
|
|
929
|
+
float | dict[str, float]
|
|
930
|
+
Weight for the specified loss term or all loss term weights
|
|
931
|
+
"""
|
|
932
|
+
loss: float | Callable[[], float]
|
|
933
|
+
key: str
|
|
934
|
+
losses: dict[str, float] = {}
|
|
935
|
+
|
|
936
|
+
if name and callable(loss := self._loss_weights[name]):
|
|
937
|
+
return loss()
|
|
938
|
+
if name:
|
|
939
|
+
return cast(float, self._loss_weights[name])
|
|
940
|
+
|
|
941
|
+
for key, loss in self._loss_weights.items():
|
|
942
|
+
losses[key] = loss() if callable(loss) else loss
|
|
943
|
+
return losses
|
|
944
|
+
|
|
945
|
+
def get_param_groups(self, learning_rate: float | tuple[float, ...] | None) -> ParamsT:
|
|
946
|
+
"""
|
|
947
|
+
Gets the parameter groups for the optimiser.
|
|
948
|
+
|
|
949
|
+
Parameters
|
|
950
|
+
----------
|
|
951
|
+
learning_rate : float | tuple[float, ...] | None
|
|
952
|
+
Learning rate or learning rates for the parameter groups
|
|
953
|
+
|
|
954
|
+
Returns
|
|
955
|
+
-------
|
|
956
|
+
ParamsT
|
|
957
|
+
Parameter groups for the optimiser
|
|
958
|
+
"""
|
|
959
|
+
return [{
|
|
960
|
+
'params': self.net.parameters(),
|
|
961
|
+
'lr': learning_rate[0] if isinstance(learning_rate, tuple) else learning_rate or 0,
|
|
962
|
+
}]
|
|
963
|
+
|
|
964
|
+
def get_rng_states(self) -> dict[str, Any]:
|
|
965
|
+
"""
|
|
966
|
+
Gets the current random number generator states.
|
|
967
|
+
|
|
968
|
+
Returns
|
|
969
|
+
-------
|
|
970
|
+
dict[str, Any]
|
|
971
|
+
Random number generator states for numpy, data loaders, pytorch, and cuda
|
|
972
|
+
"""
|
|
973
|
+
return {
|
|
974
|
+
'numpy': tuple(state.tolist() if isinstance(state, ndarray) else
|
|
975
|
+
state for state in np.random.get_state()),
|
|
976
|
+
'loaders': self._loader_states,
|
|
977
|
+
'pytorch': torch.get_rng_state(),
|
|
978
|
+
'cuda': torch.cuda.get_rng_state_all() if torch.cuda.is_available() else None,
|
|
979
|
+
}
|
|
980
|
+
|
|
981
|
+
def get_save_path(self) -> str:
|
|
982
|
+
"""
|
|
983
|
+
Gets the path to save the architecture.
|
|
984
|
+
|
|
985
|
+
Returns
|
|
986
|
+
-------
|
|
987
|
+
str
|
|
988
|
+
Path to save the architecture
|
|
989
|
+
"""
|
|
990
|
+
return self._save_path
|
|
991
|
+
|
|
992
|
+
def predict(
|
|
993
|
+
self,
|
|
994
|
+
loader: DataLoader[Any],
|
|
995
|
+
*,
|
|
996
|
+
inputs: bool = False,
|
|
997
|
+
path: str = '',
|
|
998
|
+
**kwargs: Any) -> dict[str, NDArrayLike]:
|
|
999
|
+
"""
|
|
1000
|
+
Generates predictions for the architecture and can save to a file.
|
|
1001
|
+
|
|
1002
|
+
Parameters
|
|
1003
|
+
----------
|
|
1004
|
+
loader : DataLoader[Any]
|
|
1005
|
+
Data loader to generate predictions for
|
|
1006
|
+
inputs : bool, Optional
|
|
1007
|
+
If the input data should be returned and saved, default = False
|
|
1008
|
+
path : str, Optional
|
|
1009
|
+
Path as pkl file to save the predictions if they should be saved
|
|
1010
|
+
**kwargs
|
|
1011
|
+
Optional keyword arguments to pass to batch_predict
|
|
1012
|
+
|
|
1013
|
+
Returns
|
|
1014
|
+
-------
|
|
1015
|
+
dict[str, NDArrayLike]
|
|
1016
|
+
Prediction IDs, Optional inputs, target values, and predicted values of shape (N,...)
|
|
1017
|
+
and type float for dataset of size N
|
|
1018
|
+
"""
|
|
1019
|
+
t_initial: float = time()
|
|
1020
|
+
key: str
|
|
1021
|
+
ids: tuple[str, ...] | ndarray | Tensor
|
|
1022
|
+
low_dim: list[Tensor] | list[Data[Tensor]] | list[DataList[Tensor | Data[Tensor]]]
|
|
1023
|
+
high_dim: list[Tensor] | list[Data[Tensor]] | list[DataList[Tensor | Data[Tensor]]]
|
|
1024
|
+
data: list[list[NDArrayLike | None]] = []
|
|
1025
|
+
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
|
+
}
|
|
1031
|
+
datum: NDArrayLike
|
|
1032
|
+
target: TensorLike
|
|
1033
|
+
in_data: TensorLike
|
|
1034
|
+
self.train(False)
|
|
1035
|
+
|
|
1036
|
+
if 'input_' in kwargs:
|
|
1037
|
+
warn(
|
|
1038
|
+
'input_ keyword argument is deprecated, please use inputs instead',
|
|
1039
|
+
DeprecationWarning,
|
|
1040
|
+
stacklevel=2,
|
|
1041
|
+
)
|
|
1042
|
+
inputs = kwargs.pop('input_')
|
|
1043
|
+
|
|
1044
|
+
# Generate predictions
|
|
1045
|
+
with torch.no_grad(), torch.autocast(
|
|
1046
|
+
enabled=self._half,
|
|
1047
|
+
device_type=self._device.type,
|
|
1048
|
+
dtype=torch.float32):
|
|
1049
|
+
for i, (ids, low_dim, high_dim, *_) in enumerate(loader):
|
|
1050
|
+
in_data, target = self._data_loader_translation(
|
|
1051
|
+
data_collation(low_dim, data_field=True),
|
|
1052
|
+
data_collation(high_dim, data_field=True),
|
|
1053
|
+
)
|
|
1054
|
+
data.append(cast(list[NDArrayLike | None], [
|
|
1055
|
+
ids.numpy() if isinstance(ids, Tensor) else np.array(ids),
|
|
1056
|
+
*([in_data.numpy()] if inputs else []),
|
|
1057
|
+
target.numpy(),
|
|
1058
|
+
*self.batch_predict(
|
|
1059
|
+
(in_data if isinstance(in_data, Tensor) else
|
|
1060
|
+
cast(DataList[Tensor], data_collation(
|
|
1061
|
+
cast(list[Data] | list[DataList[Tensor]], [in_data]),
|
|
1062
|
+
data_field=False,
|
|
1063
|
+
))).to(self._device),
|
|
1064
|
+
**kwargs,
|
|
1065
|
+
),
|
|
1066
|
+
]))
|
|
1067
|
+
self._predict_print(i, time() - t_initial, loader)
|
|
1068
|
+
|
|
1069
|
+
# Transforms all data and saves it to a dictionary
|
|
1070
|
+
for (key, transform), datum_ in zip(transforms.items(), zip(*data)):
|
|
1071
|
+
if datum_[0] is None:
|
|
1072
|
+
continue
|
|
1073
|
+
|
|
1074
|
+
# Concatenate values
|
|
1075
|
+
datum = data_collation(list(datum_), data_field=True)
|
|
1076
|
+
|
|
1077
|
+
if isinstance(datum, DataList) and transform:
|
|
1078
|
+
data_[key] = datum.apply(transform, back=True).numpy()
|
|
1079
|
+
data_[key] = DataList(
|
|
1080
|
+
[trans(val, back=True) for val, trans in zip(
|
|
1081
|
+
datum,
|
|
1082
|
+
transform if isinstance(transform, list) else repeat(transform),
|
|
1083
|
+
)],
|
|
1084
|
+
)
|
|
1085
|
+
elif isinstance(transform, BaseTransform):
|
|
1086
|
+
assert not isinstance(datum, DataList)
|
|
1087
|
+
data_[key] = transform(datum, back=True)
|
|
1088
|
+
else:
|
|
1089
|
+
if isinstance(transform, list):
|
|
1090
|
+
self._logger.warning(f'List of transforms requires corresponding data with key '
|
|
1091
|
+
f'({key}) to be a DataList, data will not be '
|
|
1092
|
+
f'untransformed')
|
|
1093
|
+
data_[key] = datum
|
|
1094
|
+
|
|
1095
|
+
self._save_predictions(path, data_)
|
|
1096
|
+
return data_
|
|
1097
|
+
|
|
1098
|
+
def save(self) -> None:
|
|
1099
|
+
"""
|
|
1100
|
+
Saves the architecture.
|
|
1101
|
+
"""
|
|
1102
|
+
if self._save_path:
|
|
1103
|
+
try:
|
|
1104
|
+
torch.save(self, self._save_path)
|
|
1105
|
+
except (KeyboardInterrupt, SystemExit):
|
|
1106
|
+
print('Program interrupted, finishing architecture save...')
|
|
1107
|
+
torch.save(self, self._save_path)
|
|
1108
|
+
|
|
1109
|
+
def set_save_freq(self, save_freq: int) -> None:
|
|
1110
|
+
"""
|
|
1111
|
+
Sets the frequency of saving the architecture.
|
|
1112
|
+
|
|
1113
|
+
Parameters
|
|
1114
|
+
----------
|
|
1115
|
+
save_freq : int
|
|
1116
|
+
Frequency of saving the architecture in epochs
|
|
1117
|
+
"""
|
|
1118
|
+
self._save_freq = save_freq
|
|
1119
|
+
|
|
1120
|
+
@overload
|
|
1121
|
+
def set_save_path(self, name: str, states_dir: str, *, overwrite: bool = ...) -> None:
|
|
1122
|
+
...
|
|
1123
|
+
|
|
1124
|
+
@overload
|
|
1125
|
+
def set_save_path(self, name: str, *, overwrite: bool = ...) -> None:
|
|
1126
|
+
...
|
|
1127
|
+
|
|
1128
|
+
def set_save_path(self, name: str, states_dir: str = '', *, overwrite: bool = False) -> None:
|
|
1129
|
+
"""
|
|
1130
|
+
Sets the save path for the architecture.
|
|
1131
|
+
|
|
1132
|
+
Parameters
|
|
1133
|
+
----------
|
|
1134
|
+
name : str
|
|
1135
|
+
Name to save the architecture as or full path to save the architecture
|
|
1136
|
+
states_dir : str, Optional
|
|
1137
|
+
Directory to save the architecture, if empty name is treated as full path
|
|
1138
|
+
overwrite : bool, Optional
|
|
1139
|
+
If saving can overwrite an existing save file, if False and file with the same name
|
|
1140
|
+
exists, an error will be raised, default = False
|
|
1141
|
+
"""
|
|
1142
|
+
old_save: str = os.path.basename(self._save_path)
|
|
1143
|
+
|
|
1144
|
+
if states_dir:
|
|
1145
|
+
self._save_path = utils.save_name(name, states_dir, self.net.name)
|
|
1146
|
+
else:
|
|
1147
|
+
self._save_path = name if '.pth' in name else f'{name}.pth'
|
|
1148
|
+
|
|
1149
|
+
if old_save == os.path.basename(self._save_path):
|
|
1150
|
+
return
|
|
1151
|
+
|
|
1152
|
+
if os.path.exists(self._save_path) and overwrite:
|
|
1153
|
+
self._logger.warning(f'{self._save_path} already exists and will be overwritten if '
|
|
1154
|
+
f'training continues')
|
|
1155
|
+
elif os.path.exists(self._save_path):
|
|
1156
|
+
raise FileExistsError(f'{self._save_path} already exists and overwrite is False')
|
|
1157
|
+
|
|
1158
|
+
def set_loss_weights(
|
|
1159
|
+
self,
|
|
1160
|
+
*args: float | Callable[[], float],
|
|
1161
|
+
**kwargs: float | Callable[[], float]) -> None:
|
|
1162
|
+
"""
|
|
1163
|
+
Sets the weights for the loss function terms.
|
|
1164
|
+
Loss weights passed as positional arguments are set in the order of the architecture loss
|
|
1165
|
+
weight keys, while keyword arguments are set by name and will override positional arguments.
|
|
1166
|
+
|
|
1167
|
+
Parameters
|
|
1168
|
+
----------
|
|
1169
|
+
*args, **kwargs
|
|
1170
|
+
Weights for the loss function terms
|
|
1171
|
+
"""
|
|
1172
|
+
key: str
|
|
1173
|
+
value: float | Callable[[], float]
|
|
1174
|
+
|
|
1175
|
+
if not self._loss_weights:
|
|
1176
|
+
self._logger.warning('No loss weights to set for this architecture, skipping...')
|
|
1177
|
+
return
|
|
1178
|
+
|
|
1179
|
+
for key, value in zip(self._loss_weights, args):
|
|
1180
|
+
self._loss_weights[key] = value
|
|
1181
|
+
|
|
1182
|
+
for key, value in kwargs.items():
|
|
1183
|
+
if key in self._loss_weights:
|
|
1184
|
+
self._loss_weights[key] = value
|
|
1185
|
+
else:
|
|
1186
|
+
self._logger.warning(f'Loss weight ({key}) not found in architecture loss '
|
|
1187
|
+
f'weights, skipping...')
|
|
1188
|
+
|
|
1189
|
+
def set_rng_states(self, states: dict[str, Any]) -> None:
|
|
1190
|
+
"""
|
|
1191
|
+
Sets the random number generator states.
|
|
1192
|
+
|
|
1193
|
+
Parameters
|
|
1194
|
+
----------
|
|
1195
|
+
states : dict[str, Any]
|
|
1196
|
+
Random number generator states for numpy, data loaders, pytorch, and cuda
|
|
1197
|
+
"""
|
|
1198
|
+
cuda_state: Tensor | None
|
|
1199
|
+
np.random.set_state(states.pop('numpy', np.random.get_state()))
|
|
1200
|
+
self._loader_states = states.pop('loaders', None)
|
|
1201
|
+
torch.set_rng_state(states.pop('pytorch', torch.get_rng_state()))
|
|
1202
|
+
|
|
1203
|
+
if torch.cuda.is_available() and (cuda_state := states.pop('cuda_rng', None)):
|
|
1204
|
+
torch.cuda.set_rng_state_all(cuda_state)
|
|
1205
|
+
|
|
1206
|
+
def to(self, *args: Any, **kwargs: Any) -> Self:
|
|
1207
|
+
"""
|
|
1208
|
+
Move and/or cast the parameters and buffers.
|
|
1209
|
+
|
|
1210
|
+
Parameters
|
|
1211
|
+
----------
|
|
1212
|
+
*args, **kwargs
|
|
1213
|
+
Arguments to pass to torch.Tensor.to
|
|
1214
|
+
|
|
1215
|
+
Returns
|
|
1216
|
+
-------
|
|
1217
|
+
Self
|
|
1218
|
+
The architecture with parameters and buffers moved/cast
|
|
1219
|
+
"""
|
|
1220
|
+
self.net = self.net.to(*args, **kwargs)
|
|
1221
|
+
self._param_device(self.optimiser.state, *args, **kwargs)
|
|
1222
|
+
self._device, *_ = torch._C._nn._parse_to(*args, **kwargs) # pylint: disable=protected-access
|
|
1223
|
+
|
|
1224
|
+
if isinstance(self.net, OptimizedModule):
|
|
1225
|
+
self.net._orig_mod.to(*args, **kwargs)
|
|
1226
|
+
return self
|
|
1227
|
+
|
|
1228
|
+
def train(self, train: bool) -> None:
|
|
1229
|
+
"""
|
|
1230
|
+
Changes the train/eval state of the architecture.
|
|
1231
|
+
|
|
1232
|
+
Parameters
|
|
1233
|
+
----------
|
|
1234
|
+
train : bool
|
|
1235
|
+
If the architecture should be in the train state
|
|
1236
|
+
"""
|
|
1237
|
+
self._train_state = train
|
|
1238
|
+
|
|
1239
|
+
if self._train_state:
|
|
1240
|
+
self.net.train()
|
|
1241
|
+
else:
|
|
1242
|
+
self.net.eval()
|
|
1243
|
+
|
|
1244
|
+
def training(self, epochs: int, loaders: tuple[DataLoader[Any], DataLoader[Any]]) -> None:
|
|
1245
|
+
"""
|
|
1246
|
+
Trains & validates the network for each epoch.
|
|
1247
|
+
|
|
1248
|
+
Parameters
|
|
1249
|
+
----------
|
|
1250
|
+
epochs : int
|
|
1251
|
+
Number of epochs to train the network up to
|
|
1252
|
+
loaders : tuple[DataLoader[Any], DataLoader[Any]]
|
|
1253
|
+
Train and validation data loaders
|
|
1254
|
+
"""
|
|
1255
|
+
i: int
|
|
1256
|
+
t_initial: float
|
|
1257
|
+
loss: LossCT
|
|
1258
|
+
|
|
1259
|
+
# Train for each epoch
|
|
1260
|
+
for i in range(self._epoch, epochs):
|
|
1261
|
+
t_initial = time()
|
|
1262
|
+
|
|
1263
|
+
# Train network
|
|
1264
|
+
self.train(True)
|
|
1265
|
+
self.losses[0].append(self._train_val(loaders[0]))
|
|
1266
|
+
|
|
1267
|
+
# Validate network
|
|
1268
|
+
self.train(False)
|
|
1269
|
+
self.losses[1].append(self._train_val(loaders[1]))
|
|
1270
|
+
|
|
1271
|
+
try:
|
|
1272
|
+
self._step(
|
|
1273
|
+
False,
|
|
1274
|
+
i,
|
|
1275
|
+
cast(dict, self.losses[1][-1])['total'] if isinstance(self.losses[1][-1], dict)
|
|
1276
|
+
else self.losses[1][-1],
|
|
1277
|
+
dataset=loaders[0].dataset.dataset
|
|
1278
|
+
if hasattr(loaders[0], 'dataset') and hasattr(loaders[0].dataset, 'dataset')
|
|
1279
|
+
else None,
|
|
1280
|
+
)
|
|
1281
|
+
except TypeError:
|
|
1282
|
+
warn(
|
|
1283
|
+
'_step without dataset parameter is deprecated, please update'
|
|
1284
|
+
'the method to include the dataset parameter',
|
|
1285
|
+
DeprecationWarning,
|
|
1286
|
+
stacklevel=2,
|
|
1287
|
+
)
|
|
1288
|
+
self._step(
|
|
1289
|
+
False,
|
|
1290
|
+
i,
|
|
1291
|
+
cast(dict, self.losses[1][-1])['total'] if isinstance(self.losses[1][-1], dict)
|
|
1292
|
+
else self.losses[1][-1],
|
|
1293
|
+
)
|
|
1294
|
+
|
|
1295
|
+
if self.run:
|
|
1296
|
+
self.run.log(
|
|
1297
|
+
{'Train Loss': self.losses[0][-1], 'Validation Loss': self.losses[1][-1]}
|
|
1298
|
+
if isinstance(self.losses[0][-1], float) else
|
|
1299
|
+
{f'Train {key}': val for key, val in self.losses[0][-1].items()} |
|
|
1300
|
+
{f'Validation {key}': val for key, val in self.losses[1][-1].items()}
|
|
1301
|
+
)
|
|
1302
|
+
|
|
1303
|
+
# Save training progress
|
|
1304
|
+
if loaders[0].generator:
|
|
1305
|
+
self._loader_states = (
|
|
1306
|
+
loaders[0].generator.get_state(),
|
|
1307
|
+
loaders[1].generator.get_state(),
|
|
1308
|
+
)
|
|
1309
|
+
|
|
1310
|
+
self._update_epoch()
|
|
1311
|
+
|
|
1312
|
+
if self._epoch % self._save_freq == 0 or self._epoch == epochs:
|
|
1313
|
+
self.save()
|
|
1314
|
+
|
|
1315
|
+
self._epoch_print(i, epochs, time() - t_initial)
|
|
1316
|
+
|
|
1317
|
+
self.train(False)
|
|
1318
|
+
self._plot_active = False
|
|
1319
|
+
loss = self._train_val(loaders[1])
|
|
1320
|
+
print(f"\nFinal validation loss: "
|
|
1321
|
+
f"{cast(dict, loss)['total'] if isinstance(loss, dict) else loss:.3e}")
|
|
1322
|
+
|
|
1323
|
+
|
|
1324
|
+
def load_arch(num: int | str, states_dir: str, arch_name: str, **kwargs: Any) -> BaseArchitecture:
|
|
1325
|
+
"""
|
|
1326
|
+
Loads an architecture from file.
|
|
1327
|
+
|
|
1328
|
+
Parameters
|
|
1329
|
+
----------
|
|
1330
|
+
num : int | str
|
|
1331
|
+
File number or name of the saved state
|
|
1332
|
+
states_dir : str
|
|
1333
|
+
Directory to the save files
|
|
1334
|
+
arch_name : str
|
|
1335
|
+
Name of the architecture
|
|
1336
|
+
**kwargs
|
|
1337
|
+
Optional keyword arguments to pass to torch.load
|
|
1338
|
+
|
|
1339
|
+
Returns
|
|
1340
|
+
-------
|
|
1341
|
+
BaseArchitecture
|
|
1342
|
+
Saved architecture object
|
|
1343
|
+
"""
|
|
1344
|
+
return torch.load(
|
|
1345
|
+
utils.save_name(num, states_dir, arch_name),
|
|
1346
|
+
**{'map_location': 'cpu'} | kwargs,
|
|
1347
|
+
)
|
|
1348
|
+
|
|
1349
|
+
|
|
1350
|
+
__all__ = ['BaseArchitecture', 'load_arch']
|