netloader 3.11.2__tar.gz → 3.11.4__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (41) hide show
  1. {netloader-3.11.2 → netloader-3.11.4}/PKG-INFO +1 -1
  2. {netloader-3.11.2 → netloader-3.11.4}/netloader/__init__.py +1 -1
  3. {netloader-3.11.2 → netloader-3.11.4}/netloader/architectures/base.py +2 -0
  4. {netloader-3.11.2 → netloader-3.11.4}/netloader/architectures/encoder_decoder.py +15 -1
  5. {netloader-3.11.2 → netloader-3.11.4}/netloader/architectures/flows.py +13 -5
  6. {netloader-3.11.2 → netloader-3.11.4}/netloader/models/misc.py +18 -15
  7. {netloader-3.11.2 → netloader-3.11.4}/netloader/network.py +31 -9
  8. {netloader-3.11.2 → netloader-3.11.4}/netloader.egg-info/PKG-INFO +1 -1
  9. {netloader-3.11.2 → netloader-3.11.4}/LICENSE.txt +0 -0
  10. {netloader-3.11.2 → netloader-3.11.4}/README.md +0 -0
  11. {netloader-3.11.2 → netloader-3.11.4}/netloader/architectures/__init__.py +0 -0
  12. {netloader-3.11.2 → netloader-3.11.4}/netloader/architectures/utils.py +0 -0
  13. {netloader-3.11.2 → netloader-3.11.4}/netloader/data.py +0 -0
  14. {netloader-3.11.2 → netloader-3.11.4}/netloader/layers/__init__.py +0 -0
  15. {netloader-3.11.2 → netloader-3.11.4}/netloader/layers/base.py +0 -0
  16. {netloader-3.11.2 → netloader-3.11.4}/netloader/layers/blocks.py +0 -0
  17. {netloader-3.11.2 → netloader-3.11.4}/netloader/layers/convolutional.py +0 -0
  18. {netloader-3.11.2 → netloader-3.11.4}/netloader/layers/flows.py +0 -0
  19. {netloader-3.11.2 → netloader-3.11.4}/netloader/layers/linear.py +0 -0
  20. {netloader-3.11.2 → netloader-3.11.4}/netloader/layers/misc.py +0 -0
  21. {netloader-3.11.2 → netloader-3.11.4}/netloader/layers/multi_layer.py +0 -0
  22. {netloader-3.11.2 → netloader-3.11.4}/netloader/layers/pooling.py +0 -0
  23. {netloader-3.11.2 → netloader-3.11.4}/netloader/layers/recurrent.py +0 -0
  24. {netloader-3.11.2 → netloader-3.11.4}/netloader/layers/utils.py +0 -0
  25. {netloader-3.11.2 → netloader-3.11.4}/netloader/loss_funcs.py +0 -0
  26. {netloader-3.11.2 → netloader-3.11.4}/netloader/models/__init__.py +0 -0
  27. {netloader-3.11.2 → netloader-3.11.4}/netloader/models/convnext.py +0 -0
  28. {netloader-3.11.2 → netloader-3.11.4}/netloader/networks/__init__.py +0 -0
  29. {netloader-3.11.2 → netloader-3.11.4}/netloader/schedulers.py +0 -0
  30. {netloader-3.11.2 → netloader-3.11.4}/netloader/transforms.py +0 -0
  31. {netloader-3.11.2 → netloader-3.11.4}/netloader/utils/__init__.py +0 -0
  32. {netloader-3.11.2 → netloader-3.11.4}/netloader/utils/configs.py +0 -0
  33. {netloader-3.11.2 → netloader-3.11.4}/netloader/utils/transforms.py +0 -0
  34. {netloader-3.11.2 → netloader-3.11.4}/netloader/utils/types.py +0 -0
  35. {netloader-3.11.2 → netloader-3.11.4}/netloader/utils/utils.py +0 -0
  36. {netloader-3.11.2 → netloader-3.11.4}/netloader.egg-info/SOURCES.txt +0 -0
  37. {netloader-3.11.2 → netloader-3.11.4}/netloader.egg-info/dependency_links.txt +0 -0
  38. {netloader-3.11.2 → netloader-3.11.4}/netloader.egg-info/requires.txt +0 -0
  39. {netloader-3.11.2 → netloader-3.11.4}/netloader.egg-info/top_level.txt +0 -0
  40. {netloader-3.11.2 → netloader-3.11.4}/pyproject.toml +0 -0
  41. {netloader-3.11.2 → netloader-3.11.4}/setup.cfg +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: netloader
3
- Version: 3.11.2
3
+ Version: 3.11.4
4
4
  Summary: Utility to generate and train PyTorch neural network objects from JSON files
5
5
  Author-email: Ethan Tregidga <ethan.tregidga@epfl.ch>
6
6
  License-Expression: MIT
@@ -6,7 +6,7 @@ import logging
6
6
  import warnings
7
7
 
8
8
 
9
- __version__ = '3.11.2'
9
+ __version__ = '3.11.4'
10
10
  __author__ = 'Ethan Tregidga'
11
11
  logging.basicConfig(format='%(levelname)s: %(message)s', level=logging.WARNING)
12
12
  warnings.filterwarnings('once', category=DeprecationWarning, module=r'^netloader(\.|$)')
@@ -62,6 +62,8 @@ class BaseArchitecture(UtilityMixin, ABC, Generic[LossCT, TensorLossCT]):
62
62
  Architecture optimiser
63
63
  scheduler : LRScheduler
64
64
  Optimiser scheduler
65
+ run : Run | None
66
+ Weights & Biases run object for logging and tracking experiments
65
67
  net : BaseNetwork
66
68
  Neural network
67
69
  """
@@ -27,7 +27,7 @@ class BaseEncoder(BaseArchitecture):
27
27
  description : str
28
28
  Description of the architecture
29
29
  version : str
30
- Version of the architecture
30
+ Version of the architecture when it was created or re-saved
31
31
  losses : tuple[list[LossCT], list[LossCT]]
32
32
  Architecture training and validation losses as a float or dictionary of losses for each loss
33
33
  function
@@ -42,6 +42,8 @@ class BaseEncoder(BaseArchitecture):
42
42
  Architecture optimiser
43
43
  scheduler : LRScheduler
44
44
  Optimiser scheduler
45
+ run : Run | None
46
+ Weights & Biases run object for logging and tracking experiments
45
47
  net : BaseNetwork
46
48
  Neural network
47
49
  """
@@ -211,6 +213,8 @@ class Autoencoder(BaseArchitecture):
211
213
  ----------
212
214
  description : str
213
215
  Description of the architecture
216
+ version : str
217
+ Version of the architecture when it was created or re-saved
214
218
  losses : tuple[list[LossCT], list[LossCT]]
215
219
  Architecture training and validation losses as a float or dictionary of losses for each loss
216
220
  function
@@ -223,6 +227,8 @@ class Autoencoder(BaseArchitecture):
223
227
  Architecture optimiser
224
228
  scheduler : LRScheduler
225
229
  Optimiser scheduler
230
+ run : Run | None
231
+ Weights & Biases run object for logging and tracking experiments
226
232
  net : BaseNetwork
227
233
  Neural network
228
234
  reconstruct_func : BaseLoss
@@ -427,6 +433,8 @@ class Decoder(BaseArchitecture):
427
433
  ----------
428
434
  description : str
429
435
  Description of the architecture
436
+ version : str
437
+ Version of the architecture when it was created or re-saved
430
438
  losses : tuple[list[LossCT], list[LossCT]]
431
439
  Architecture training and validation losses as a float or dictionary of losses for each loss
432
440
  function
@@ -441,6 +449,8 @@ class Decoder(BaseArchitecture):
441
449
  Architecture optimiser
442
450
  scheduler : LRScheduler
443
451
  Optimiser scheduler
452
+ run : Run | None
453
+ Weights & Biases run object for logging and tracking experiments
444
454
  net : BaseNetwork
445
455
  Neural network
446
456
  """
@@ -587,6 +597,8 @@ class Encoder(BaseEncoder):
587
597
  ----------
588
598
  description : str
589
599
  Description of the architecture
600
+ version : str
601
+ Version of the architecture when it was created or re-saved
590
602
  losses : tuple[list[LossCT], list[LossCT]]
591
603
  Architecture training and validation losses as a float or dictionary of losses for each loss
592
604
  function
@@ -601,6 +613,8 @@ class Encoder(BaseEncoder):
601
613
  Architecture optimiser
602
614
  scheduler : LRScheduler
603
615
  Optimiser scheduler
616
+ run : Run | None
617
+ Weights & Biases run object for logging and tracking experiments
604
618
  net : BaseNetwork
605
619
  Neural network
606
620
  """
@@ -14,8 +14,8 @@ from zuko.distributions import NormalizingFlow
14
14
  from netloader.data import Data
15
15
  from netloader.utils import label_change
16
16
  from netloader.loss_funcs import BaseLoss
17
- from netloader.network import BaseNetwork
18
17
  from netloader.transforms import BaseTransform
18
+ from netloader.network import BaseNetwork, Network
19
19
  from netloader.architectures.base import BaseArchitecture
20
20
  from netloader.architectures.encoder_decoder import BaseEncoder
21
21
  from netloader.utils.types import NDArrayLike, TensorListLike, NDArrayListLike, TensorT
@@ -29,10 +29,10 @@ class NormFlow(BaseArchitecture):
29
29
 
30
30
  Attributes
31
31
  ----------
32
- net : BaseNetwork
33
- Neural spline flow
34
32
  description : str
35
33
  Description of the architecture
34
+ version : str
35
+ Version of the architecture when it was created or re-saved
36
36
  losses : tuple[list[LossCT], list[LossCT]]
37
37
  Architecture training and validation losses as a float or dictionary of losses for each loss
38
38
  function
@@ -45,6 +45,10 @@ class NormFlow(BaseArchitecture):
45
45
  Architecture optimiser
46
46
  scheduler : LRScheduler
47
47
  Optimiser scheduler
48
+ run : Run | None
49
+ Weights & Biases run object for logging and tracking experiments
50
+ net : BaseNetwork
51
+ Neural spline flow
48
52
  """
49
53
  @staticmethod
50
54
  def _data_loader_translation(low_dim: TensorT, high_dim: TensorT) -> tuple[TensorT, TensorT]:
@@ -132,6 +136,8 @@ class NormFlowEncoder(BaseEncoder):
132
136
  ----------
133
137
  description : str
134
138
  Description of the architecture
139
+ version : str
140
+ Version of the architecture when it was created or re-saved
135
141
  losses : tuple[list[LossCT], list[LossCT]]
136
142
  Architecture training and validation losses as a float or dictionary of losses for each loss
137
143
  function
@@ -146,6 +152,8 @@ class NormFlowEncoder(BaseEncoder):
146
152
  Architecture optimiser
147
153
  scheduler : LRScheduler
148
154
  Optimiser scheduler
155
+ run : Run | None
156
+ Weights & Biases run object for logging and tracking experiments
149
157
  net : BaseNetwork
150
158
  Neural network
151
159
  """
@@ -222,7 +230,7 @@ class NormFlowEncoder(BaseEncoder):
222
230
  net,
223
231
  overwrite=overwrite,
224
232
  mix_precision=mix_precision,
225
- learning_rate=0,
233
+ learning_rate=0.,
226
234
  description=description,
227
235
  verbose=verbose,
228
236
  classes=classes,
@@ -243,7 +251,7 @@ class NormFlowEncoder(BaseEncoder):
243
251
  'max': transform,
244
252
  'meds': transform,
245
253
  }
246
- assert isinstance(self.net, nn.ModuleList)
254
+ assert isinstance(self.net, Network)
247
255
 
248
256
  if not self._train_encoder:
249
257
  self.net.layers[:-1].requires_grad_(False)
@@ -12,7 +12,7 @@ from netloader.utils import Config, TypedModuleList
12
12
  from netloader.network import BaseNetwork, CompatibleNetwork, Network
13
13
 
14
14
 
15
- class MultiNetwork(BaseNetwork[TypedModuleList[CompatibleNetwork | Network]]):
15
+ class MultiNetwork(BaseNetwork[TypedModuleList[BaseNetwork]]):
16
16
  """
17
17
  Network class to hold multiple networks and train them together.
18
18
 
@@ -27,7 +27,7 @@ class MultiNetwork(BaseNetwork[TypedModuleList[CompatibleNetwork | Network]]):
27
27
  is the batch size
28
28
  kl_loss : Tensor
29
29
  KL divergence loss on the latent space of shape (1) and type float, if using a sample layer
30
- layers : TypedModuleList[CompatibleNetwork | Network]
30
+ layers : TypedModuleList[BaseNetwork]
31
31
  Network construction
32
32
  """
33
33
  _cache: bool
@@ -36,14 +36,14 @@ class MultiNetwork(BaseNetwork[TypedModuleList[CompatibleNetwork | Network]]):
36
36
  def __init__(
37
37
  self,
38
38
  name: str,
39
- *nets: nn.Module | nn.ModuleList | CompatibleNetwork | Network,
39
+ *nets: nn.Module | nn.ModuleList | BaseNetwork,
40
40
  save_outputs: bool = ...) -> None: ...
41
41
 
42
42
  @overload
43
43
  def __init__(
44
44
  self,
45
45
  name: str,
46
- *nets: str | nn.Module | nn.ModuleList | CompatibleNetwork | Network | Config,
46
+ *nets: str | nn.Module | nn.ModuleList | BaseNetwork | Config,
47
47
  save_outputs: bool = ...,
48
48
  suppress_warning: bool = ...,
49
49
  root: str = ...,
@@ -55,7 +55,7 @@ class MultiNetwork(BaseNetwork[TypedModuleList[CompatibleNetwork | Network]]):
55
55
  def __init__(
56
56
  self,
57
57
  name: str,
58
- *nets: str | nn.Module | nn.ModuleList | CompatibleNetwork | Network | Config,
58
+ *nets: str | nn.Module | nn.ModuleList | BaseNetwork | Config,
59
59
  save_outputs: bool = False,
60
60
  suppress_warning: bool = False,
61
61
  root: str = '',
@@ -68,7 +68,7 @@ class MultiNetwork(BaseNetwork[TypedModuleList[CompatibleNetwork | Network]]):
68
68
  ----------
69
69
  name : str
70
70
  Name of the network, used for saving
71
- *nets: str | nn.Module | nn.ModuleList | CompatibleNetwork | Network | Config
71
+ *nets: str | nn.Module | nn.ModuleList | BaseNetwork | Config
72
72
  Networks to combine sequentially
73
73
  save_outputs : bool, Optional
74
74
  If outputs from each network should be saved, default = False
@@ -91,7 +91,7 @@ class MultiNetwork(BaseNetwork[TypedModuleList[CompatibleNetwork | Network]]):
91
91
  self.name = name
92
92
  self.layers = TypedModuleList()
93
93
  calc_in_shape: bool = True
94
- net: str | nn.Module | nn.ModuleList | CompatibleNetwork | Network | Config
94
+ net: str | nn.Module | nn.ModuleList | BaseNetwork | Config
95
95
  builds: list[bool] = [isinstance(net, (str, Config)) for net in nets]
96
96
  out_shape: list[int] | None
97
97
  default: dict[str, Any] | None
@@ -150,15 +150,15 @@ class MultiNetwork(BaseNetwork[TypedModuleList[CompatibleNetwork | Network]]):
150
150
  in_shape = out_shape
151
151
 
152
152
  @overload
153
- def __getitem__(self, item: int) -> CompatibleNetwork | Network: ...
153
+ def __getitem__(self, item: int) -> BaseNetwork: ...
154
154
 
155
155
  @overload
156
- def __getitem__(self, item: slice) -> TypedModuleList[CompatibleNetwork | Network]: ...
156
+ def __getitem__(self, item: slice) -> TypedModuleList[BaseNetwork]: ...
157
157
 
158
158
  def __getitem__(
159
159
  self,
160
160
  item: int | slice,
161
- ) -> CompatibleNetwork | Network | TypedModuleList[CompatibleNetwork | Network]:
161
+ ) -> BaseNetwork | TypedModuleList[BaseNetwork]:
162
162
  """
163
163
  Returns the network(s) at the given index or slice.
164
164
 
@@ -169,7 +169,7 @@ class MultiNetwork(BaseNetwork[TypedModuleList[CompatibleNetwork | Network]]):
169
169
 
170
170
  Returns
171
171
  -------
172
- CompatibleNetwork | Network | TypedModuleList[CompatibleNetwork | Network]
172
+ BaseNetwork | TypedModuleList[BaseNetwork]
173
173
  The network(s) at the given index or slice
174
174
  """
175
175
  return self.layers[item]
@@ -183,7 +183,7 @@ class MultiNetwork(BaseNetwork[TypedModuleList[CompatibleNetwork | Network]]):
183
183
  def __setstate__(self, state: dict[str, Any]) -> None:
184
184
  super().__setstate__(state)
185
185
  self._cache = state.get('save_outputs', False)
186
- self.layers = TypedModuleList(state.get('layers', []))
186
+ self.layers = TypedModuleList(state.get('layers', state.get('net')))
187
187
 
188
188
  def forward(self, x: TensorListLike) -> TensorListLike:
189
189
  """
@@ -200,7 +200,7 @@ class MultiNetwork(BaseNetwork[TypedModuleList[CompatibleNetwork | Network]]):
200
200
  TensorListLike
201
201
  Output tensor from the network with shape (N,...) and type float
202
202
  """
203
- net: CompatibleNetwork | Network
203
+ net: BaseNetwork
204
204
  self.checkpoints = []
205
205
 
206
206
  for i, net in enumerate(self.layers):
@@ -231,11 +231,14 @@ class MultiNetwork(BaseNetwork[TypedModuleList[CompatibleNetwork | Network]]):
231
231
  dict[str, Any] | Config
232
232
  Network configuration
233
233
  """
234
- net: CompatibleNetwork | Network
234
+ net: BaseNetwork
235
235
  config = Config(layers=[net.get_config() for net in self.layers])
236
236
  return config.to_dict() if dict_ else config
237
237
 
238
238
  def to(self, *args: Any, **kwargs: Any) -> Self:
239
239
  super().to(*args, **kwargs)
240
- self.layers.to(*args, **kwargs)
240
+ layer: BaseNetwork
241
+
242
+ for layer in self.layers:
243
+ layer.to(*args, **kwargs)
241
244
  return self
@@ -81,7 +81,7 @@ class BaseNetwork(nn.Module, ABC, Generic[ModuleListT]):
81
81
  """
82
82
  if item == 'net':
83
83
  warn(
84
- 'CompatibleNetwork.net is deprecated, please use CompatibleNetwork.layers '
84
+ 'BaseNetwork.net is deprecated, please use BaseNetwork.layers '
85
85
  'instead',
86
86
  DeprecationWarning,
87
87
  stacklevel=2,
@@ -100,6 +100,28 @@ class BaseNetwork(nn.Module, ABC, Generic[ModuleListT]):
100
100
  """
101
101
  return {'name': self.name, 'version': netloader.__version__}
102
102
 
103
+ def __setattr__(self, key: str, value: Any) -> None:
104
+ """
105
+ Sets the attribute of the network.
106
+
107
+ Parameters
108
+ ----------
109
+ key : str
110
+ Name of the attribute to set
111
+ value : Any
112
+ Value to set the attribute to
113
+ """
114
+ if key == 'net':
115
+ warn(
116
+ 'BaseNetwork.net is deprecated, please use BaseNetwork.layers '
117
+ 'instead',
118
+ DeprecationWarning,
119
+ stacklevel=2,
120
+ )
121
+ super().__setattr__('layers', value)
122
+ else:
123
+ super().__setattr__(key, value)
124
+
103
125
  def __setstate__(self, state: dict[str, Any]) -> None:
104
126
  """
105
127
  Sets the state of the network for pickling
@@ -181,7 +203,14 @@ class BaseNetwork(nn.Module, ABC, Generic[ModuleListT]):
181
203
 
182
204
  def to(self, *args: Any, **kwargs: Any) -> Self:
183
205
  super().to(*args, **kwargs)
184
- self.kl_loss = self.kl_loss.to(*args, **kwargs)
206
+ self.kl_loss = self.kl_loss.detach().to(*args, **kwargs)
207
+
208
+ for i, checkpoint in enumerate(self.checkpoints):
209
+ if hasattr(checkpoint, 'to'):
210
+ self.checkpoints[i] = checkpoint.detach().to(*args, **kwargs)
211
+ else:
212
+ self._logger.warning(f'Failed to move checkpoint {checkpoint.__class__.__name__} '
213
+ f'to device')
185
214
  return self
186
215
 
187
216
 
@@ -522,13 +551,6 @@ class Network(BaseNetwork[TypedModuleList[layers.BaseLayer]]):
522
551
 
523
552
  for layer in self.layers:
524
553
  layer.to(*args, **kwargs)
525
-
526
- for i, checkpoint in enumerate(self.checkpoints):
527
- if hasattr(checkpoint, 'to'):
528
- self.checkpoints[i] = checkpoint.to(*args, **kwargs)
529
- else:
530
- self._logger.warning(f'Failed to move checkpoint {checkpoint.__class__.__name__} '
531
- f'to device')
532
554
  return self
533
555
 
534
556
 
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: netloader
3
- Version: 3.11.2
3
+ Version: 3.11.4
4
4
  Summary: Utility to generate and train PyTorch neural network objects from JSON files
5
5
  Author-email: Ethan Tregidga <ethan.tregidga@epfl.ch>
6
6
  License-Expression: MIT
File without changes
File without changes
File without changes
File without changes
File without changes