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,123 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Normalising flow network layers
|
|
3
|
+
"""
|
|
4
|
+
from typing import Any
|
|
5
|
+
|
|
6
|
+
from torch import Tensor
|
|
7
|
+
from zuko.flows import NSF
|
|
8
|
+
from zuko.distributions import NormalizingFlow
|
|
9
|
+
|
|
10
|
+
from netloader.utils import Shapes
|
|
11
|
+
from netloader.layers.base import BaseLayer
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class SplineFlow(BaseLayer):
|
|
15
|
+
"""
|
|
16
|
+
Neural spline flow layer constructor.
|
|
17
|
+
|
|
18
|
+
Should only be used as the last layer in the network as it does not return a Tensor.
|
|
19
|
+
|
|
20
|
+
Attributes
|
|
21
|
+
----------
|
|
22
|
+
group : int
|
|
23
|
+
Layer group, if 0 it will always be used, else it will only be used if its group matches the
|
|
24
|
+
Networks
|
|
25
|
+
description : str
|
|
26
|
+
Description of the layer
|
|
27
|
+
"""
|
|
28
|
+
def __init__(
|
|
29
|
+
self,
|
|
30
|
+
transforms: int,
|
|
31
|
+
hidden_features: list[int],
|
|
32
|
+
net_out: list[int],
|
|
33
|
+
shapes: Shapes,
|
|
34
|
+
*,
|
|
35
|
+
context: bool = False,
|
|
36
|
+
features: int | None = None,
|
|
37
|
+
factor: float | None = None,
|
|
38
|
+
**kwargs: Any) -> None:
|
|
39
|
+
"""
|
|
40
|
+
Generates a neural spline flow (NSF) for use in BaseNetwork
|
|
41
|
+
|
|
42
|
+
Adds attributes of name ('flow'), optimiser (Adam), and scheduler (ReduceLROnPlateau)
|
|
43
|
+
|
|
44
|
+
Parameters
|
|
45
|
+
----------
|
|
46
|
+
transforms : int
|
|
47
|
+
Number of transforms
|
|
48
|
+
hidden_features : list[int]
|
|
49
|
+
Number of features in each of the hidden layers
|
|
50
|
+
net_out : list[int]
|
|
51
|
+
Shape of the network's output
|
|
52
|
+
shapes : Shapes
|
|
53
|
+
Shape of the outputs from each layer
|
|
54
|
+
context : bool, Optional
|
|
55
|
+
If the output from the previous layer should be used to condition the normalising flow,
|
|
56
|
+
default = False
|
|
57
|
+
features : int | None, Optional
|
|
58
|
+
Dimensions of the probability distribution, if factor is provided, features will not be
|
|
59
|
+
used
|
|
60
|
+
factor : float | None, Optional
|
|
61
|
+
Output features is equal to the factor of the network's output, will be used if
|
|
62
|
+
provided, else features will be used
|
|
63
|
+
**kwargs
|
|
64
|
+
Leftover parameters to pass to base layer for checking
|
|
65
|
+
"""
|
|
66
|
+
super().__init__(**kwargs)
|
|
67
|
+
self._context: bool = context
|
|
68
|
+
self._transforms: int = transforms
|
|
69
|
+
self._hidden_features: list[int] = hidden_features
|
|
70
|
+
self._layer: NSF
|
|
71
|
+
context_: int
|
|
72
|
+
shape: list[int] = shapes[-1].copy()
|
|
73
|
+
|
|
74
|
+
# Number of features can be defined by either a factor of the output size or explicitly
|
|
75
|
+
if factor:
|
|
76
|
+
shape[-1] = max(1, int(net_out[-1] * factor))
|
|
77
|
+
elif features:
|
|
78
|
+
shape[-1] = features
|
|
79
|
+
else:
|
|
80
|
+
raise ValueError('Either features or factor must be provided and be non-zero')
|
|
81
|
+
|
|
82
|
+
if self._context:
|
|
83
|
+
context_ = shapes[-1][-1]
|
|
84
|
+
else:
|
|
85
|
+
context_ = 0
|
|
86
|
+
|
|
87
|
+
self._layer = NSF(
|
|
88
|
+
features=shape[-1],
|
|
89
|
+
context=context_,
|
|
90
|
+
transforms=transforms,
|
|
91
|
+
hidden_features=hidden_features,
|
|
92
|
+
)
|
|
93
|
+
|
|
94
|
+
shapes.append(shape)
|
|
95
|
+
|
|
96
|
+
def __getstate__(self) -> dict[str, Any]:
|
|
97
|
+
return super().__getstate__() | {
|
|
98
|
+
'context': self._context,
|
|
99
|
+
'transforms': self._transforms,
|
|
100
|
+
'features': self._layer.transform.transforms[0].passes,
|
|
101
|
+
'hidden_features': self._hidden_features,
|
|
102
|
+
}
|
|
103
|
+
|
|
104
|
+
def forward(self, x: Tensor, *_: Any, **__: Any) -> NormalizingFlow:
|
|
105
|
+
"""
|
|
106
|
+
Forward pass of the neural spline flow layer
|
|
107
|
+
|
|
108
|
+
Parameters
|
|
109
|
+
----------
|
|
110
|
+
x : Tensor
|
|
111
|
+
Input tensor with shape (N,...) and type float, where N is the batch size
|
|
112
|
+
|
|
113
|
+
Returns
|
|
114
|
+
-------
|
|
115
|
+
NormalizingFlow
|
|
116
|
+
Normalising flow distribution
|
|
117
|
+
"""
|
|
118
|
+
if self._context:
|
|
119
|
+
return self._layer(x)
|
|
120
|
+
return self._layer()
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
__all__ = ['SplineFlow']
|
|
@@ -0,0 +1,445 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Linear network layers
|
|
3
|
+
"""
|
|
4
|
+
from __future__ import annotations
|
|
5
|
+
import logging as log
|
|
6
|
+
from inspect import signature
|
|
7
|
+
from typing import TYPE_CHECKING, Any, Self, Literal, cast
|
|
8
|
+
|
|
9
|
+
import torch
|
|
10
|
+
import numpy as np
|
|
11
|
+
from torch import nn, Tensor
|
|
12
|
+
|
|
13
|
+
from netloader.utils import Shapes, compare_versions
|
|
14
|
+
from netloader.layers.base import BaseLayer, BaseSingleLayer
|
|
15
|
+
|
|
16
|
+
if TYPE_CHECKING:
|
|
17
|
+
from netloader.network import Network
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class Activation(BaseSingleLayer):
|
|
21
|
+
"""
|
|
22
|
+
Activation layer constructor.
|
|
23
|
+
|
|
24
|
+
Attributes
|
|
25
|
+
----------
|
|
26
|
+
group : int
|
|
27
|
+
Layer group, if 0 it will always be used, else it will only be used if its group matches the
|
|
28
|
+
Networks
|
|
29
|
+
description : str
|
|
30
|
+
Description of the layer
|
|
31
|
+
layers : Sequential
|
|
32
|
+
Layers to loop through in the forward pass
|
|
33
|
+
"""
|
|
34
|
+
def __init__(
|
|
35
|
+
self,
|
|
36
|
+
*,
|
|
37
|
+
activation: str = 'ELU',
|
|
38
|
+
shapes: Shapes | None = None,
|
|
39
|
+
activation_kwargs: dict[str, Any] | None = None,
|
|
40
|
+
**kwargs: Any) -> None:
|
|
41
|
+
"""
|
|
42
|
+
Parameters
|
|
43
|
+
----------
|
|
44
|
+
activation : str, Optional
|
|
45
|
+
Which activation function to use from PyTorch, default = 'ELU'
|
|
46
|
+
shapes : Shapes | None, Optional
|
|
47
|
+
Shape of the outputs from each layer, only required if tracking layer outputs is
|
|
48
|
+
necessary
|
|
49
|
+
activation_kwargs : dict[str, Any] | None, Optional
|
|
50
|
+
Additional keyword arguments to pass to the activation function
|
|
51
|
+
**kwargs
|
|
52
|
+
Leftover parameters to pass to base layer for checking
|
|
53
|
+
"""
|
|
54
|
+
super().__init__(**({'idx': 0} | kwargs))
|
|
55
|
+
self._kwargs: dict[str, Any] = activation_kwargs or {}
|
|
56
|
+
self.layers.add_module('Activation', getattr(nn, activation)(
|
|
57
|
+
**({'inplace': True} if 'inplace' in signature(getattr(nn, activation)).parameters
|
|
58
|
+
else {}) | self._kwargs,
|
|
59
|
+
))
|
|
60
|
+
|
|
61
|
+
# If not used as a layer in a network
|
|
62
|
+
if not shapes:
|
|
63
|
+
return
|
|
64
|
+
|
|
65
|
+
shapes.append(shapes[-1].copy())
|
|
66
|
+
|
|
67
|
+
def __getstate__(self) -> dict[str, Any]:
|
|
68
|
+
return super().__getstate__() | {
|
|
69
|
+
'activation': self.layers.Activation.__class__.__name__,
|
|
70
|
+
'activation_kwargs': self._kwargs,
|
|
71
|
+
}
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
class Linear(BaseSingleLayer):
|
|
75
|
+
"""
|
|
76
|
+
Linear layer constructor.
|
|
77
|
+
|
|
78
|
+
Attributes
|
|
79
|
+
----------
|
|
80
|
+
group : int
|
|
81
|
+
Layer group, if 0 it will always be used, else it will only be used if its group matches the
|
|
82
|
+
Networks
|
|
83
|
+
description : str
|
|
84
|
+
Description of the layer
|
|
85
|
+
layers : Sequential
|
|
86
|
+
Layers to loop through in the forward pass
|
|
87
|
+
"""
|
|
88
|
+
def __init__(
|
|
89
|
+
self,
|
|
90
|
+
net_out: list[int],
|
|
91
|
+
shapes: Shapes,
|
|
92
|
+
*,
|
|
93
|
+
features: int | None = None,
|
|
94
|
+
layer: int | None = None,
|
|
95
|
+
factor: float | None = None,
|
|
96
|
+
batch_norm: bool = False,
|
|
97
|
+
flatten_target: bool = False,
|
|
98
|
+
dropout: float = 0,
|
|
99
|
+
activation: str | None = 'SELU',
|
|
100
|
+
**kwargs: Any) -> None:
|
|
101
|
+
"""
|
|
102
|
+
Parameters
|
|
103
|
+
----------
|
|
104
|
+
net_out : list[int]
|
|
105
|
+
Shape of the network's output
|
|
106
|
+
shapes : Shapes
|
|
107
|
+
Shape of the outputs from each layer
|
|
108
|
+
features : int | None, Optional
|
|
109
|
+
Number of output features for the layer, if factor is provided, features will not be
|
|
110
|
+
used
|
|
111
|
+
layer : int | None, Optional
|
|
112
|
+
If factor is not None, which layer for factor to be relative to, if None, network output
|
|
113
|
+
will be used
|
|
114
|
+
factor : float | None, Optional
|
|
115
|
+
Output features is equal to the factor of the network's output, or if layer is provided,
|
|
116
|
+
which layer to be relative to, will be used if provided, else features will be used
|
|
117
|
+
batch_norm : bool, Optional
|
|
118
|
+
If batch normalisation should be used, default = False
|
|
119
|
+
flatten_target : bool, Optional
|
|
120
|
+
If the target should be flattened so that features is equal to the product of the target
|
|
121
|
+
multiplied by factor, if factor is provided, default = False
|
|
122
|
+
dropout : float, Optional
|
|
123
|
+
Probability of dropout, default = 0
|
|
124
|
+
activation : str | None, Optional
|
|
125
|
+
Which activation function to use from PyTorch, default = 'SELU'
|
|
126
|
+
**kwargs
|
|
127
|
+
Leftover parameters to pass to base layer for checking
|
|
128
|
+
"""
|
|
129
|
+
super().__init__(**kwargs)
|
|
130
|
+
shape: list[int] = shapes[-1].copy()
|
|
131
|
+
target: list[int] = shapes[layer] if layer is not None else net_out
|
|
132
|
+
|
|
133
|
+
# Number of features can be defined by either a factor of the output size or explicitly
|
|
134
|
+
shape = self._check_factor_filters(
|
|
135
|
+
shape[::-1],
|
|
136
|
+
filters=features,
|
|
137
|
+
factor=factor,
|
|
138
|
+
target=[int(np.prod(target))] if flatten_target else target[::-1],
|
|
139
|
+
)[::-1]
|
|
140
|
+
|
|
141
|
+
self.layers.add_module('Linear', nn.Linear(
|
|
142
|
+
in_features=shapes[-1][-1],
|
|
143
|
+
out_features=shape[-1],
|
|
144
|
+
))
|
|
145
|
+
|
|
146
|
+
# Optional layers
|
|
147
|
+
if activation:
|
|
148
|
+
self.layers.add_module('Activation', getattr(nn, activation)(
|
|
149
|
+
**{'inplace': True} if 'inplace' in signature(getattr(nn, activation)).parameters
|
|
150
|
+
else {},
|
|
151
|
+
))
|
|
152
|
+
|
|
153
|
+
if batch_norm:
|
|
154
|
+
self.layers.add_module('BatchNorm', nn.BatchNorm1d(shape[0]))
|
|
155
|
+
|
|
156
|
+
if dropout:
|
|
157
|
+
self.layers.add_module('Dropout', nn.Dropout(dropout))
|
|
158
|
+
|
|
159
|
+
shapes.append(shape)
|
|
160
|
+
|
|
161
|
+
def __getstate__(self) -> dict[str, Any]:
|
|
162
|
+
return super().__getstate__() | {
|
|
163
|
+
'features': self.layers.Linear.out_features, # type: ignore[union-attr]
|
|
164
|
+
'batch_norm': 'BatchNorm' in self.layers._modules,
|
|
165
|
+
'dropout': self.layers.Dropout.p if 'Dropout' in self.layers._modules else 0, # type: ignore[union-attr]
|
|
166
|
+
'activation': self.layers.Activation.__class__.__name__
|
|
167
|
+
if 'Activation' in self.layers._modules else None,
|
|
168
|
+
}
|
|
169
|
+
|
|
170
|
+
|
|
171
|
+
class OrderedBottleneck(BaseLayer):
|
|
172
|
+
"""
|
|
173
|
+
Information-ordered bottleneck to randomly change the size of the bottleneck in an autoencoder
|
|
174
|
+
to encode the most important information in the first values of the latent space.
|
|
175
|
+
|
|
176
|
+
See `Information-Ordered Bottlenecks for Adaptive Semantic Compression
|
|
177
|
+
<https://arxiv.org/abs/2305.11213>`_ by Ho et al. (2023).
|
|
178
|
+
|
|
179
|
+
Attributes
|
|
180
|
+
----------
|
|
181
|
+
group : int
|
|
182
|
+
Layer group, if 0 it will always be used, else it will only be used if its group matches the
|
|
183
|
+
Networks
|
|
184
|
+
min_size : int
|
|
185
|
+
Minimum gate size
|
|
186
|
+
description : str
|
|
187
|
+
Description of the layer
|
|
188
|
+
"""
|
|
189
|
+
def __init__(self, shapes: Shapes, *, min_size: int = 0, **kwargs: Any) -> None:
|
|
190
|
+
"""
|
|
191
|
+
Parameters
|
|
192
|
+
----------
|
|
193
|
+
shapes : Shapes
|
|
194
|
+
Shape of the outputs from each layer
|
|
195
|
+
min_size : int, Optional
|
|
196
|
+
Minimum gate size, default = 0
|
|
197
|
+
**kwargs
|
|
198
|
+
Leftover parameters to pass to base layer for checking
|
|
199
|
+
"""
|
|
200
|
+
super().__init__(**kwargs)
|
|
201
|
+
self._offset: bool = compare_versions(self._ver, '3.10.0')
|
|
202
|
+
self.min_size: int = min_size
|
|
203
|
+
shapes.append(shapes[-1].copy())
|
|
204
|
+
|
|
205
|
+
def __getstate__(self) -> dict[str, Any]:
|
|
206
|
+
return super().__getstate__() | {'min_size': self.min_size}
|
|
207
|
+
|
|
208
|
+
def forward(self, x: Tensor, *_: Any, **__: Any) -> Tensor:
|
|
209
|
+
"""
|
|
210
|
+
Forward pass of the information-ordered bottleneck layer
|
|
211
|
+
|
|
212
|
+
Parameters
|
|
213
|
+
----------
|
|
214
|
+
x : Tensor
|
|
215
|
+
Input tensor with shape (N,...,Z) and type float, where N is the batch size and Z is the
|
|
216
|
+
latent space where Z > min_size
|
|
217
|
+
|
|
218
|
+
Returns
|
|
219
|
+
-------
|
|
220
|
+
Tensor
|
|
221
|
+
Input zeroed from a random index to the last value along the dimension Z with shape
|
|
222
|
+
(N,...,Z) and type float
|
|
223
|
+
"""
|
|
224
|
+
if not self.training or self.min_size >= x.size(-1):
|
|
225
|
+
return x
|
|
226
|
+
|
|
227
|
+
idx: Tensor = torch.randint(
|
|
228
|
+
self.min_size,
|
|
229
|
+
x.size(-1) + int(self._offset),
|
|
230
|
+
(1,),
|
|
231
|
+
device=x.device,
|
|
232
|
+
)
|
|
233
|
+
gate: Tensor = torch.where(
|
|
234
|
+
torch.arange(x.size(-1), device=x.device) < idx + int(not self._offset),
|
|
235
|
+
torch.ones_like(idx),
|
|
236
|
+
torch.zeros_like(idx),
|
|
237
|
+
)
|
|
238
|
+
return x * gate
|
|
239
|
+
|
|
240
|
+
def extra_repr(self) -> str:
|
|
241
|
+
"""
|
|
242
|
+
Displays layer parameters when printing the network
|
|
243
|
+
|
|
244
|
+
Returns
|
|
245
|
+
-------
|
|
246
|
+
str
|
|
247
|
+
Layer parameters
|
|
248
|
+
"""
|
|
249
|
+
return f'min_size={self.min_size}'
|
|
250
|
+
|
|
251
|
+
|
|
252
|
+
class Sample(BaseLayer):
|
|
253
|
+
"""
|
|
254
|
+
Samples random values from a Gaussian distribution for a variational autoencoder,
|
|
255
|
+
mean and standard deviation are the first and second half of the input channels with the last
|
|
256
|
+
channel ignored if there are an odd number.
|
|
257
|
+
|
|
258
|
+
Attributes
|
|
259
|
+
----------
|
|
260
|
+
group : int
|
|
261
|
+
Layer group, if 0 it will always be used, else it will only be used if its group matches the
|
|
262
|
+
Networks
|
|
263
|
+
description : str
|
|
264
|
+
Description of the layer
|
|
265
|
+
sample_layer : Normal
|
|
266
|
+
Layer to sample values from a Gaussian distribution
|
|
267
|
+
"""
|
|
268
|
+
def __init__(self, idx: int, shapes: Shapes, **kwargs: Any) -> None:
|
|
269
|
+
"""
|
|
270
|
+
Parameters
|
|
271
|
+
----------
|
|
272
|
+
idx : int
|
|
273
|
+
Layer number
|
|
274
|
+
shapes : Shapes
|
|
275
|
+
Shape of the outputs from each layer
|
|
276
|
+
**kwargs
|
|
277
|
+
Leftover parameters to pass to base layer for checking
|
|
278
|
+
"""
|
|
279
|
+
super().__init__(idx=idx, **kwargs)
|
|
280
|
+
self.sample_layer: torch.distributions.Normal = torch.distributions.Normal(
|
|
281
|
+
torch.tensor(0.).to(self._device),
|
|
282
|
+
torch.tensor(1.).to(self._device),
|
|
283
|
+
)
|
|
284
|
+
|
|
285
|
+
if shapes[-1][0] % 2 == 1:
|
|
286
|
+
log.warning(f'Sample in layer {idx} expects an even length along the first dimension, '
|
|
287
|
+
f'but input shape is {shapes[-1]}, last element will be ignored')
|
|
288
|
+
|
|
289
|
+
shapes.append(shapes[-1].copy())
|
|
290
|
+
shapes[-1][0] = shapes[-1][0] // 2
|
|
291
|
+
|
|
292
|
+
def forward(self, x: Tensor, net: Network, *_: Any, **__: Any) -> Tensor:
|
|
293
|
+
"""
|
|
294
|
+
Forward pass of the sampling layer for a variational autoencoder
|
|
295
|
+
|
|
296
|
+
Parameters
|
|
297
|
+
----------
|
|
298
|
+
x : Tensor
|
|
299
|
+
Input tensor with shape (N,C,...) | (N,Z) and type float, where N is the batch size, and
|
|
300
|
+
either the channels dimension, C, or latent, Z, containing the mean and standard
|
|
301
|
+
deviation
|
|
302
|
+
net : Network
|
|
303
|
+
Parent network that this layer is part of
|
|
304
|
+
|
|
305
|
+
Returns
|
|
306
|
+
-------
|
|
307
|
+
Tensor
|
|
308
|
+
Output tensor sampled from the input tensor split into mean and standard deviation with
|
|
309
|
+
shape (N,C/2,...) | (N,Z/2) and type float
|
|
310
|
+
"""
|
|
311
|
+
split: int = x.size(1) // 2
|
|
312
|
+
mean: Tensor = x[:, :split]
|
|
313
|
+
std: Tensor = torch.exp(x[:, split:2 * split])
|
|
314
|
+
x = mean + std * self.sample_layer.rsample(mean.shape)
|
|
315
|
+
|
|
316
|
+
net.kl_loss = 0.5 * torch.mean(
|
|
317
|
+
mean ** 2 + std ** 2 - 2 * torch.log(std) - 1
|
|
318
|
+
)
|
|
319
|
+
return x
|
|
320
|
+
|
|
321
|
+
def to(self, *args: Any, **kwargs: Any) -> Self:
|
|
322
|
+
super().to(*args, **kwargs)
|
|
323
|
+
self.sample_layer.loc = self.sample_layer.loc.to(*args, **kwargs)
|
|
324
|
+
self.sample_layer.scale = self.sample_layer.scale.to(*args, **kwargs)
|
|
325
|
+
return self
|
|
326
|
+
|
|
327
|
+
|
|
328
|
+
class Upsample(BaseSingleLayer):
|
|
329
|
+
"""
|
|
330
|
+
Constructs an upsampler.
|
|
331
|
+
|
|
332
|
+
Attributes
|
|
333
|
+
----------
|
|
334
|
+
group : int
|
|
335
|
+
Layer group, if 0 it will always be used, else it will only be used if its group matches the
|
|
336
|
+
Networks
|
|
337
|
+
description : str
|
|
338
|
+
Description of the layer
|
|
339
|
+
layers : Sequential
|
|
340
|
+
Layers to loop through in the forward pass
|
|
341
|
+
"""
|
|
342
|
+
def __init__(
|
|
343
|
+
self,
|
|
344
|
+
idx: int,
|
|
345
|
+
shapes: Shapes,
|
|
346
|
+
*,
|
|
347
|
+
shape: list[int] | None = None,
|
|
348
|
+
scale: float | list[float] | tuple[float, ...] = 2,
|
|
349
|
+
mode: Literal['nearest', 'linear', 'bilinear', 'bicubic', 'trilinear'] = 'nearest',
|
|
350
|
+
**kwargs: Any) -> None:
|
|
351
|
+
"""
|
|
352
|
+
Parameters
|
|
353
|
+
----------
|
|
354
|
+
idx : int
|
|
355
|
+
Layer number
|
|
356
|
+
shapes : Shapes
|
|
357
|
+
Shape of the outputs from each layer
|
|
358
|
+
shape : list[int] | None, Optional
|
|
359
|
+
Shape of the output, will be used if provided, else scale will be used
|
|
360
|
+
scale : float | tuple[float] | tuple[float, ...], Optional
|
|
361
|
+
Factor to upscale all or individual dimensions, first dimension is ignored, won't be
|
|
362
|
+
used if shape is provided, default = 2
|
|
363
|
+
mode : {'nearest', 'linear', 'bilinear', 'bicubic', 'trilinear'}
|
|
364
|
+
What interpolation method to use for upsampling
|
|
365
|
+
**kwargs
|
|
366
|
+
Leftover parameters to pass to base layer for checking
|
|
367
|
+
"""
|
|
368
|
+
super().__init__(idx=idx, **kwargs)
|
|
369
|
+
modes: dict[str, list[int]] = {
|
|
370
|
+
'nearest': [2, 3, 4],
|
|
371
|
+
'linear': [2],
|
|
372
|
+
'bilinear': [3],
|
|
373
|
+
'bicubic': [3],
|
|
374
|
+
'trilinear': [4],
|
|
375
|
+
}
|
|
376
|
+
|
|
377
|
+
# Check for errors
|
|
378
|
+
self._check_shape((2, 4), shapes[-1])
|
|
379
|
+
self._check_upsample(shapes[-1], shape)
|
|
380
|
+
self._check_options('mode', mode, set(modes))
|
|
381
|
+
self._check_mode_dimension(mode, shapes[-1], modes)
|
|
382
|
+
|
|
383
|
+
if isinstance(scale, list):
|
|
384
|
+
scale = tuple(scale)
|
|
385
|
+
elif not isinstance(scale, tuple):
|
|
386
|
+
scale = (scale,) * len(shapes[-1][1:])
|
|
387
|
+
|
|
388
|
+
if shape:
|
|
389
|
+
shapes.append(shape)
|
|
390
|
+
else:
|
|
391
|
+
shapes.append(shapes[-1].copy())
|
|
392
|
+
shapes[-1][1:] = [
|
|
393
|
+
int(length * factor) for length, factor in zip(shapes[-1][1:], scale)
|
|
394
|
+
]
|
|
395
|
+
|
|
396
|
+
self.layers.add_module('Upsample', nn.Upsample(
|
|
397
|
+
size=tuple(shape) if shape else None,
|
|
398
|
+
scale_factor=None if shape else scale,
|
|
399
|
+
mode=mode,
|
|
400
|
+
))
|
|
401
|
+
|
|
402
|
+
def __getstate__(self) -> dict[str, Any]:
|
|
403
|
+
layer: nn.Upsample = cast(nn.Upsample, self.layers.Upsample)
|
|
404
|
+
return super().__getstate__() | {
|
|
405
|
+
'shape': list(shape) if isinstance(shape := layer.size, tuple) else shape,
|
|
406
|
+
'scale': layer.scale_factor,
|
|
407
|
+
'mode': layer.mode,
|
|
408
|
+
}
|
|
409
|
+
|
|
410
|
+
@staticmethod
|
|
411
|
+
def _check_mode_dimension(mode: str, shape: list[int], modes: dict[str, list[int]]) -> None:
|
|
412
|
+
"""
|
|
413
|
+
Checks if the upsampling mode supports the number of input dimensions
|
|
414
|
+
|
|
415
|
+
Parameters
|
|
416
|
+
----------
|
|
417
|
+
mode : str
|
|
418
|
+
Current upsampling mode
|
|
419
|
+
shape : list[int]
|
|
420
|
+
Input shape
|
|
421
|
+
modes : dict[str, list[int]]
|
|
422
|
+
Modes with a list of supported dimensions
|
|
423
|
+
"""
|
|
424
|
+
if len(shape) not in modes[mode]:
|
|
425
|
+
raise ValueError(f'{mode} only supports input dimensions of {modes[mode]}, '
|
|
426
|
+
f'input shape is {shape}')
|
|
427
|
+
|
|
428
|
+
@staticmethod
|
|
429
|
+
def _check_upsample(in_shape: list[int], out_shape: list[int] | None) -> None:
|
|
430
|
+
"""
|
|
431
|
+
Checks if the input shape has the same number of dimensions as the output shape
|
|
432
|
+
|
|
433
|
+
Parameters
|
|
434
|
+
----------
|
|
435
|
+
in_shape : list[int]
|
|
436
|
+
Input shape
|
|
437
|
+
out_shape : list[int] | None
|
|
438
|
+
Target output shape
|
|
439
|
+
"""
|
|
440
|
+
if out_shape is not None and len(out_shape) != 1 and len(out_shape) + 1 != len(in_shape):
|
|
441
|
+
raise ValueError(f'Target shape {out_shape} is not compatible with input shape '
|
|
442
|
+
f'{in_shape}, check that channels dimension is not in target shape')
|
|
443
|
+
|
|
444
|
+
|
|
445
|
+
__all__ = ['Activation', 'Linear', 'OrderedBottleneck', 'Sample', 'Upsample']
|