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.
@@ -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']