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,1008 @@
1
+ """
2
+ Convolutional network layers
3
+ """
4
+ from inspect import signature
5
+ from typing import Any, Type, Literal
6
+
7
+ import torch
8
+ import numpy as np
9
+ from torch import nn, Tensor
10
+
11
+ from netloader.layers.misc import LayerNorm, Pad
12
+ from netloader.utils import Shapes, compare_versions
13
+ from netloader.layers.base import BaseLayer, BaseSingleLayer
14
+ from netloader.layers.utils import _int_list_conversion, _kernel_shape, _padding
15
+
16
+
17
+ class Conv(BaseSingleLayer):
18
+ """
19
+ Convolutional layer constructor.
20
+
21
+ Supports 1D, 2D, and 3D convolution.
22
+
23
+ Attributes
24
+ ----------
25
+ group : int
26
+ Layer group, if 0 it will always be used, else it will only be used if its group matches the
27
+ Networks
28
+ description : str
29
+ Description of the layer
30
+ layers : Sequential
31
+ Layers to loop through in the forward pass
32
+ """
33
+ def __init__(
34
+ self,
35
+ net_out: list[int],
36
+ shapes: Shapes,
37
+ *,
38
+ filters: int | None = None,
39
+ layer: int | None = None,
40
+ factor: float | None = None,
41
+ groups: int = 1,
42
+ kernel: int | list[int] = 3,
43
+ stride: int | list[int] = 1,
44
+ padding: int | Literal['same'] | list[int] = 0,
45
+ dropout: float = 0,
46
+ activation: str | None = 'ELU',
47
+ norm: Literal[None, 'batch', 'layer'] = None,
48
+ padding_mode: Literal[None, 'replicate', 'zeros', 'reflect', 'circular'] = None,
49
+ **kwargs: Any) -> None:
50
+ """
51
+ Parameters
52
+ ----------
53
+ net_out : list[int]
54
+ Shape of the network's output, required only if layer contains factor
55
+ shapes : Shapes
56
+ Shape of the outputs from each layer
57
+ filters : int | None, Optional
58
+ Number of convolutional filters, will be used if provided, else factor will be used
59
+ layer : int | None, Optional
60
+ If factor is not None, which layer for factor to be relative to, if None, network output
61
+ will be used
62
+ factor : float | None, Optional
63
+ Number of convolutional filters equal to the output channels, or if layer is provided,
64
+ the layer's channels, multiplied by factor, won't be used if filters is provided
65
+ groups : int, Optional
66
+ Number of input channel groups, each with its own convolutional filter(s), input and
67
+ output channels must both be divisible by the number of groups, default = 1
68
+ kernel : int | list[int], Optional
69
+ Size of the kernel, default = 3
70
+ stride : int | list[int], Optional
71
+ Stride of the kernel, default = 1
72
+ padding : int | {'same'} | list[int], Optional
73
+ Input padding, can an int, list of ints or 'same' where 'same' preserves the input
74
+ shape, default = 0
75
+ dropout : float, Optional
76
+ Probability of dropout, default = 0
77
+ activation : str | None, Optional
78
+ Which activation function to use from PyTorch, default = 'ELU'
79
+ norm : {None, 'batch', 'layer'}
80
+ If batch or layer normalisation should be used
81
+ padding_mode : {None, 'zeros', 'reflect', 'replicate', 'circular'}
82
+ Padding mode to use from PyTorch, if None zero padding is used if version is >3.9.4 else
83
+ replication padding is used
84
+ **kwargs
85
+ Leftover parameters to pass to base layer for checking
86
+ """
87
+ super().__init__(**kwargs)
88
+ self._pad: int | Literal['same'] | list[int] = padding
89
+ asymmetric: bool = False
90
+ padding_: int | list[int]
91
+ shape: list[int] = shapes[-1].copy()
92
+ target: list[int] = shapes[layer] if layer is not None else net_out
93
+ conv: Type[nn.Module]
94
+ dropout_: Type[nn.Module]
95
+ batch_norm_: Type[nn.Module]
96
+
97
+ # Check for errors and calculate same padding
98
+ if isinstance(self._pad, str):
99
+ self._check_options('padding', self._pad, {'same'})
100
+ self._check_stride(stride)
101
+ asymmetric, padding_ = _padding(kernel, stride, shape, shape)
102
+ else:
103
+ padding_ = self._pad
104
+
105
+ if padding_mode is None:
106
+ padding_mode = 'zeros' if compare_versions(self._ver, '3.10.0') else \
107
+ 'replicate'
108
+
109
+ if asymmetric and compare_versions(self._ver, '3.10.1'):
110
+ assert isinstance(padding_, list)
111
+ self.layers.add_module('Pad', Pad(
112
+ tuple(padding_),
113
+ mode='constant' if padding_mode == 'zeros' else padding_mode,
114
+ ))
115
+ self._pad = 0
116
+
117
+ # Check for errors and calculate number of filters
118
+ self._check_shape((2, 4), shape)
119
+ self._check_factor_filters(shape, filters=filters, factor=factor, target=target)
120
+ self._check_groups(shapes[-1][0], shape[0], groups)
121
+ self._check_options('norm', norm, {None, 'batch', 'layer'})
122
+
123
+ conv, dropout_, batch_norm_ = [
124
+ [nn.Conv1d, nn.Dropout1d, nn.BatchNorm1d],
125
+ [nn.Conv2d, nn.Dropout2d, nn.BatchNorm2d],
126
+ [nn.Conv3d, nn.Dropout3d, nn.BatchNorm3d],
127
+ ][len(shape) - 2]
128
+
129
+ self.layers.add_module('Conv', conv(
130
+ in_channels=shapes[-1][0],
131
+ out_channels=shape[0],
132
+ kernel_size=kernel,
133
+ stride=stride,
134
+ padding=self._pad,
135
+ groups=groups,
136
+ padding_mode=padding_mode,
137
+ ))
138
+
139
+ # Optional layers
140
+ if activation:
141
+ self.layers.add_module('Activation', getattr(nn, activation)(
142
+ **{'inplace': True} if 'inplace' in signature(getattr(nn, activation)).parameters
143
+ else {},
144
+ ))
145
+
146
+ if norm == 'batch':
147
+ self.layers.add_module('BatchNorm', batch_norm_(shape[0]))
148
+ elif norm == 'layer':
149
+ self.layers.add_module('LayerNorm', LayerNorm(shape=shape[0:1]))
150
+
151
+ if dropout:
152
+ self.layers.add_module('Dropout', dropout_(dropout))
153
+
154
+ if self._pad == 'same' or asymmetric:
155
+ shapes.append(shape)
156
+ else:
157
+ shapes.append(_kernel_shape(kernel, stride, padding_, shape))
158
+
159
+ def __getstate__(self) -> dict[str, Any]:
160
+ return super().__getstate__() | {
161
+ 'filters': self.layers.Conv.out_channels, # type: ignore[union-attr]
162
+ 'groups': self.layers.Conv.groups, # type: ignore[union-attr]
163
+ 'kernel': self.layers.Conv.kernel_size, # type: ignore[union-attr]
164
+ 'stride': self.layers.Conv.stride,
165
+ 'padding': self._pad,
166
+ 'dropout': self.layers.Dropout.p if hasattr(self.layers, 'Dropout') else 0, # type: ignore[union-attr]
167
+ 'activation': self.layers.Activation.__class__.__name__
168
+ if hasattr(self.layers, 'Activation') else None,
169
+ 'norm': 'batch' if hasattr(self.layers, 'BatchNorm') else
170
+ 'layer' if hasattr(self.layers, 'LayerNorm') else None,
171
+ 'padding_mode': self.layers.Conv.padding_mode, # type: ignore[union-attr]
172
+ }
173
+
174
+ @staticmethod
175
+ def _check_groups(in_channels: int, out_channels: int, groups: int) -> None:
176
+ """
177
+ Checks if the number of groups is compatible with the number of input and output channels
178
+
179
+ Parameters
180
+ ----------
181
+ in_channels : int
182
+ Number of input channels
183
+ out_channels : int
184
+ Number of output channels
185
+ groups : int
186
+ Number of groups
187
+ """
188
+ if in_channels % groups != 0 or out_channels % groups != 0:
189
+ raise ValueError(f'Number of groups ({groups}) is not compatible with input channels '
190
+ f'({in_channels}) and/or output channels ({out_channels}), check that '
191
+ f'they are divisible by groups')
192
+
193
+
194
+ class ConvDepth(Conv):
195
+ """
196
+ Constructs a depthwise convolutional layer.
197
+
198
+ Supports 1D, 2D, and 3D convolution.
199
+
200
+ Attributes
201
+ ----------
202
+ group : int
203
+ Layer group, if 0 it will always be used, else it will only be used if its group matches the
204
+ Networks
205
+ description : str
206
+ Description of the layer
207
+ layers : Sequential
208
+ Layers to loop through in the forward pass
209
+ """
210
+ def __init__(
211
+ self,
212
+ net_out: list[int],
213
+ shapes: Shapes,
214
+ *,
215
+ filters: int | None = None,
216
+ layer: int | None = None,
217
+ factor: float | None = None,
218
+ kernel: int | list[int] = 3,
219
+ stride: int | list[int] = 1,
220
+ padding: int | Literal['same'] | list[int] = 0,
221
+ dropout: float = 0,
222
+ activation: str | None = 'ELU',
223
+ norm: Literal[None, 'batch', 'layer'] = None,
224
+ padding_mode: Literal[None, 'zeros', 'reflect', 'replicate', 'circular'] = None,
225
+ **kwargs: Any) -> None:
226
+ """
227
+ Parameters
228
+ ----------
229
+ net_out : list[int]
230
+ Shape of the network's output, required only if layer contains factor
231
+ shapes : Shapes
232
+ Shape of the outputs from each layer
233
+ filters : int | None, Optional
234
+ Number of convolutional filters, will be used if provided, else factor will be used
235
+ layer : int | None, Optional
236
+ If factor is not None, which layer for factor to be relative to, if None, network output
237
+ will be used
238
+ factor : float | None, Optional
239
+ Number of convolutional filters equal to the output channels multiplied by factor,
240
+ won't be used if filters is provided
241
+ kernel : int | list[int], Optional
242
+ Size of the kernel, default = 3
243
+ stride : int | list[int], Optional
244
+ Stride of the kernel, default = 1
245
+ padding : int | {'same'} | list[int], Optional
246
+ Input padding, can an int, list of ints or 'same' where 'same' preserves the input
247
+ shape, default = 0
248
+ dropout : float, Optional
249
+ Probability of dropout, default = 0
250
+ activation : str | None, Optional
251
+ Which activation function to use, default = 'ELU'
252
+ norm : {None, 'batch', 'layer'}
253
+ If batch or layer normalisation should be used
254
+ padding_mode : {None, 'zeros', 'reflect', 'replicate', 'circular'}
255
+ Padding mode to use from PyTorch, if None zero padding is used if version is >3.9.4 else
256
+ replication padding is used
257
+ **kwargs
258
+ Leftover parameters to pass to base layer for checking
259
+ """
260
+ super().__init__(
261
+ net_out=net_out,
262
+ shapes=shapes,
263
+ filters=filters,
264
+ layer=layer,
265
+ factor=factor,
266
+ groups=shapes[-1][0],
267
+ kernel=kernel,
268
+ stride=stride,
269
+ padding=padding,
270
+ dropout=dropout,
271
+ activation=activation,
272
+ norm=norm,
273
+ padding_mode=padding_mode,
274
+ **kwargs,
275
+ )
276
+
277
+ def __getstate__(self) -> dict[str, Any]:
278
+ state: dict[str, Any] = super().__getstate__()
279
+ state.pop('groups', None)
280
+ return state
281
+
282
+
283
+ class ConvDepthDownscale(Conv):
284
+ """
285
+ Constructs depth downscaler using convolution with kernel size of 1.
286
+
287
+ Attributes
288
+ ----------
289
+ group : int
290
+ Layer group, if 0 it will always be used, else it will only be used if its group matches the
291
+ Networks
292
+ description : str
293
+ Description of the layer
294
+ layers : Sequential
295
+ Layers to loop through in the forward pass
296
+ """
297
+ def __init__(
298
+ self,
299
+ net_out: list[int],
300
+ shapes: Shapes,
301
+ *,
302
+ dropout: float = 0,
303
+ activation: str | None = 'ELU',
304
+ norm: Literal[None, 'batch', 'layer'] = None,
305
+ padding_mode: Literal[None, 'zeros', 'reflect', 'replicate', 'circular'] = None,
306
+ **kwargs: Any) -> None:
307
+ """
308
+ Parameters
309
+ ----------
310
+ net_out : list[int]
311
+ Shape of the network's output, required only if layer contains factor
312
+ shapes: Shapes
313
+ Shape of the outputs from each layer
314
+ dropout : float, Optional
315
+ Probability of dropout, default = 0
316
+ activation : str | None, Optional
317
+ Which activation function to use, default = 'ELU'
318
+ norm : {None, 'batch', 'layer'}
319
+ If batch or layer normalisation should be used
320
+ padding_mode : {None, 'zeros', 'reflect', 'replicate', 'circular'}
321
+ Padding mode to use from PyTorch, if None zero padding is used if version is >3.9.4 else
322
+ replication padding is used
323
+ **kwargs
324
+ Leftover parameters to pass to base layer for checking
325
+ """
326
+ super().__init__(
327
+ net_out=net_out,
328
+ shapes=shapes,
329
+ filters=1,
330
+ stride=1,
331
+ kernel=1,
332
+ padding='same',
333
+ dropout=dropout,
334
+ activation=activation,
335
+ norm=norm,
336
+ padding_mode=padding_mode,
337
+ **kwargs,
338
+ )
339
+
340
+ def __getstate__(self) -> dict[str, Any]:
341
+ state: dict[str, Any] = super().__getstate__()
342
+ state.pop('filters', None)
343
+ state.pop('stride', None)
344
+ state.pop('kernel', None)
345
+ state.pop('padding', None)
346
+ return state
347
+
348
+
349
+ class ConvDownscale(Conv):
350
+ """
351
+ Constructs a strided convolutional layer for downscaling.
352
+
353
+ The scale factor is equal to the stride and kernel size.
354
+
355
+ Attributes
356
+ ----------
357
+ group : int
358
+ Layer group, if 0 it will always be used, else it will only be used if its group matches the
359
+ Networks
360
+ description : str
361
+ Description of the layer
362
+ layers : Sequential
363
+ Layers to loop through in the forward pass
364
+ """
365
+ def __init__(self,
366
+ net_out: list[int],
367
+ shapes: Shapes,
368
+ *,
369
+ filters: int | None = None,
370
+ layer: int | None = None,
371
+ factor: float | None = None,
372
+ scale: int = 2,
373
+ dropout: float = 0,
374
+ activation: str | None = 'ELU',
375
+ norm: Literal[None, 'batch', 'layer'] = None,
376
+ **kwargs: Any) -> None:
377
+ """
378
+ Parameters
379
+ ----------
380
+ net_out : list[int]
381
+ Shape of the network's output, required only if layer contains factor
382
+ shapes: Shapes
383
+ Shape of the outputs from each layer
384
+ filters : int | None, Optional
385
+ Number of convolutional filters, will be used if provided, else factor will be used
386
+ layer : int | None, Optional
387
+ If factor is not None, which layer for factor to be relative to, if None, network output
388
+ will be used
389
+ factor : float | None, Optional
390
+ Number of convolutional filters equal to the output channels multiplied by factor,
391
+ won't be used if filters is provided
392
+ scale : int, Optional
393
+ Stride and size of the kernel, which acts as the downscaling factor, default = 2
394
+ dropout : float, Optional
395
+ Probability of dropout, default = 0
396
+ activation : str | None, Optional
397
+ Which activation function to use, default = 'ELU'
398
+ norm : {None, 'batch', 'layer'}
399
+ If batch or layer normalisation should be used
400
+ **kwargs
401
+ Leftover parameters to pass to base layer for checking
402
+ """
403
+ super().__init__(
404
+ net_out=net_out,
405
+ shapes=shapes,
406
+ filters=filters,
407
+ layer=layer,
408
+ factor=factor,
409
+ kernel=scale,
410
+ stride=scale,
411
+ padding=0,
412
+ dropout=dropout,
413
+ activation=activation,
414
+ norm=norm,
415
+ **kwargs,
416
+ )
417
+
418
+ def __getstate__(self) -> dict[str, Any]:
419
+ state: dict[str, Any] = super().__getstate__()
420
+ state.pop('padding', None)
421
+ state['scale'] = state.pop('stride')
422
+ state.pop('kernel', None)
423
+ return state
424
+
425
+
426
+ class ConvTranspose(BaseSingleLayer):
427
+ """
428
+ Constructs a transpose convolutional layer with fractional stride for input upscaling.
429
+
430
+ Supports 1D, 2D, and 3D transposed convolution.
431
+
432
+ Attributes
433
+ ----------
434
+ group : int
435
+ Layer group, if 0 it will always be used, else it will only be used if its group matches the
436
+ Networks
437
+ description : str
438
+ Description of the layer
439
+ layers : Sequential
440
+ Layers to loop through in the forward pass
441
+ """
442
+ def __init__(
443
+ self,
444
+ net_out: list[int],
445
+ shapes: Shapes,
446
+ *,
447
+ filters: int | None = None,
448
+ layer: int | None = None,
449
+ factor: float | None = None,
450
+ kernel: int | list[int] = 3,
451
+ stride: int | list[int] = 1,
452
+ out_padding: int | list[int] = 0,
453
+ dilation: int | list[int] = 1,
454
+ padding: int | Literal['same'] | list[int] = 0,
455
+ dropout: float = 0,
456
+ padding_mode: Literal['zeros', 'reflect', 'replicate', 'circular'] = 'zeros',
457
+ activation: str | None = 'ELU',
458
+ norm: Literal[None, 'batch', 'layer'] = None,
459
+ **kwargs: Any) -> None:
460
+ """
461
+ Parameters
462
+ ----------
463
+ net_out : list[int]
464
+ Shape of the network's output, required only if layer contains factor
465
+ shapes: Shapes
466
+ Shape of the outputs from each layer
467
+ filters : int | None, Optional
468
+ Number of convolutional filters, will be used if provided, else factor will be used
469
+ layer : int | None, Optional
470
+ If factor is not None, which layer for factor to be relative to, if None, network output
471
+ will be used
472
+ factor : float | None, Optional
473
+ Number of convolutional filters equal to the output channels multiplied by factor,
474
+ won't be used if filters is provided
475
+ kernel : int | list[int], Optional
476
+ Size of the kernel, default = 3
477
+ stride : int | list[int], Optional
478
+ Stride of the kernel, default = 1
479
+ out_padding : int | list[int], Optional
480
+ Padding applied to the output, default = 0
481
+ dilation : int | list[int], Optional
482
+ Spacing between kernel points, default = 1
483
+ padding : int | {'same'} | list[int], Optional
484
+ Inverse of convolutional padding which removes rows from each dimension in the output,
485
+ default = 0
486
+ dropout : float, Optional
487
+ Probability of dropout, default = 0
488
+ padding_mode : {'zeros', 'reflect', 'replicate', 'circular'}
489
+ Padding mode to use from PyTorch
490
+ activation : str | None, Optional
491
+ Which activation function to use, default = 'ELU'
492
+ norm : {None, 'batch', 'layer'}
493
+ If batch or layer normalisation should be used
494
+ **kwargs
495
+ Leftover parameters to pass to base layer for checking
496
+ """
497
+ super().__init__(**kwargs)
498
+ self._slice: np.ndarray = np.array([slice(None)] * len(shapes[-1][1:]))
499
+ padding_: int | str | list[int] = padding
500
+ shape: list[int] = shapes[-1].copy()
501
+ target: list[int] = shapes[layer] if layer is not None else net_out
502
+ transpose: Type[nn.Module]
503
+ dropout_: Type[nn.Module]
504
+ batch_norm_: Type[nn.Module]
505
+
506
+ # Check for errors and calculate same padding
507
+ if isinstance(padding, str):
508
+ self._check_options('padding', padding, {'same'})
509
+ padding = _padding_transpose(kernel, stride, dilation, shapes[-1], shape)
510
+
511
+ # Check for errors and calculate number of filters
512
+ self._check_shape((2, 4), shape)
513
+ self._check_out_padding(stride, dilation, out_padding)
514
+ self._check_factor_filters(shape, filters=filters, factor=factor, target=target)
515
+ self._check_options('norm', norm, {None, 'batch', 'layer'})
516
+
517
+ transpose, dropout_, batch_norm_ = [
518
+ [nn.ConvTranspose1d, nn.Dropout1d, nn.BatchNorm1d],
519
+ [nn.ConvTranspose2d, nn.Dropout2d, nn.BatchNorm2d],
520
+ [nn.ConvTranspose3d, nn.Dropout3d, nn.BatchNorm3d],
521
+ ][len(shape) - 2]
522
+
523
+ shape = _kernel_transpose_shape(
524
+ kernel,
525
+ stride,
526
+ padding,
527
+ dilation,
528
+ out_padding,
529
+ shape,
530
+ )
531
+
532
+ # Correct same padding for one-sided padding
533
+ if padding_ == 'same' and shape != shapes[-1]:
534
+ self._slice[np.array(shape[1:]) - np.array(shapes[-1][1:]) == 1] = slice(-1)
535
+ shape[1:] = shapes[-1][1:]
536
+
537
+ self.layers.add_module('Transpose', transpose(
538
+ in_channels=shapes[-1][0],
539
+ out_channels=shape[0],
540
+ kernel_size=kernel,
541
+ stride=stride,
542
+ padding=padding,
543
+ output_padding=out_padding,
544
+ padding_mode=padding_mode,
545
+ dilation=dilation,
546
+ ))
547
+
548
+ # Optional layers
549
+ if activation:
550
+ self.layers.add_module('Activation', getattr(nn, activation)(
551
+ **{'inplace': True} if 'inplace' in signature(getattr(nn, activation)).parameters
552
+ else {},
553
+ ))
554
+
555
+ if norm == 'batch':
556
+ self.layers.add_module('BatchNorm', batch_norm_(shape[0]))
557
+ elif norm == 'layer':
558
+ self.layers.add_module('LayerNorm', LayerNorm(shape=shape[0:1]))
559
+
560
+ if dropout:
561
+ self.layers.add_module('Dropout', dropout_(dropout))
562
+
563
+ shapes.append(shape)
564
+
565
+ def __getstate__(self) -> dict[str, Any]:
566
+ return super().__getstate__() | {
567
+ 'filters': self.layers.Transpose.out_channels, # type: ignore[union-attr]
568
+ 'kernel': self.layers.Transpose.kernel_size, # type: ignore[union-attr]
569
+ 'stride': self.layers.Transpose.stride,
570
+ 'out_padding': self.layers.Transpose.output_padding, # type: ignore[union-attr]
571
+ 'dilation': self.layers.Transpose.dilation, # type: ignore[union-attr]
572
+ 'padding': self.layers.Transpose.padding, # type: ignore[union-attr]
573
+ 'dropout': self.layers.Dropout.p if hasattr(self.layers, 'Dropout') else 0, # type: ignore[union-attr]
574
+ 'padding_mode': self.layers.Transpose.padding_mode, # type: ignore[union-attr]
575
+ 'activation': self.layers.Activation.__class__.__name__
576
+ if hasattr(self.layers, 'Activation') else None,
577
+ 'norm': 'batch' if hasattr(self.layers, 'BatchNorm') else
578
+ 'layer' if hasattr(self.layers, 'LayerNorm') else None,
579
+ }
580
+
581
+ @staticmethod
582
+ def _check_out_padding(
583
+ stride: int | list[int],
584
+ dilation: int | list[int],
585
+ out_padding: int | list[int]) -> None:
586
+ """
587
+ Checks if the output padding is compatible with the dilation and stride
588
+
589
+ Parameters
590
+ ----------
591
+ stride : int | list[int]
592
+ Stride of the kernel
593
+ dilation : int | list[int]
594
+ Dilation of the kernel
595
+ out_padding : int | list[int]
596
+ Output padding
597
+ """
598
+ if ((np.array(out_padding) >= np.array(stride)) *
599
+ (np.array(out_padding) >= np.array(dilation))).any():
600
+ raise ValueError(f'Output padding ({out_padding}) must be smaller than either stride '
601
+ f'({stride}) or dilation ({dilation})')
602
+
603
+ def forward(self, x: Tensor, *_: Any, **__: Any) -> Tensor:
604
+ """
605
+ Forward pass of the transposed convolutional layer
606
+
607
+ Parameters
608
+ ----------
609
+ x : Tensor
610
+ Input tensor with shape (N,...) and type float, where N is the batch size
611
+
612
+ Returns
613
+ -------
614
+ Tensor
615
+ Output tensor with shape (N,...) and type float
616
+ """
617
+ x = super().forward(x)
618
+ return x[..., *self._slice]
619
+
620
+
621
+ class ConvTransposeUpscale(ConvTranspose):
622
+ """
623
+ Constructs an upscaler using a transposed convolutional layer.
624
+
625
+ Supports 1D, 2D, and 3D transposed convolutional upscaling.
626
+
627
+ Attributes
628
+ ----------
629
+ group : int
630
+ Layer group, if 0 it will always be used, else it will only be used if its group matches the
631
+ Networks
632
+ description : str
633
+ Description of the layer
634
+ layers : Sequential
635
+ Layers to loop through in the forward pass
636
+ """
637
+ def __init__(
638
+ self,
639
+ net_out: list[int],
640
+ shapes: Shapes,
641
+ *,
642
+ filters: int | None = None,
643
+ layer: int | None = None,
644
+ factor: float | None = None,
645
+ scale: int | list[int] = 2,
646
+ out_padding: int | list[int] = 0,
647
+ dropout: float = 0,
648
+ activation: str | None = 'ELU',
649
+ norm: Literal[None, 'batch', 'layer'] = None,
650
+ **kwargs: Any) -> None:
651
+ """
652
+ Parameters
653
+ ----------
654
+ net_out : list[int]
655
+ Shape of the network's output, required only if layer contains factor
656
+ shapes: Shapes
657
+ Shape of the outputs from each layer
658
+ filters : int | None, Optional
659
+ Number of convolutional filters, will be used if provided, else factor will be used
660
+ layer : int | None, Optional
661
+ If factor is not None, which layer for factor to be relative to, if None, network output
662
+ will be used
663
+ factor : float | None, Optional
664
+ Number of convolutional filters equal to the output channels multiplied by factor,
665
+ won't be used if filters is provided
666
+ scale : int | list[int], Optional
667
+ Stride and size of the kernel, which acts as the upscaling factor, default = 2
668
+ out_padding : int | list[int], Optional
669
+ Padding applied to the output, default = 0
670
+ dropout : float, Optional
671
+ Probability of dropout, default = 0
672
+ activation : str | None, Optional
673
+ Which activation function to use, default = 'ELU'
674
+ norm : {None, 'batch', 'layer'}
675
+ If batch or layer normalisation should be used
676
+ **kwargs
677
+ Leftover parameters to pass to base layer for checking
678
+ """
679
+ super().__init__(
680
+ net_out=net_out,
681
+ shapes=shapes,
682
+ filters=filters,
683
+ layer=layer,
684
+ factor=factor,
685
+ kernel=scale,
686
+ stride=scale,
687
+ padding=0,
688
+ out_padding=out_padding,
689
+ dropout=dropout,
690
+ activation=activation,
691
+ norm=norm,
692
+ **kwargs,
693
+ )
694
+
695
+ def __getstate__(self) -> dict[str, Any]:
696
+ state: dict[str, Any] = super().__getstate__()
697
+ state.pop('padding', None)
698
+ return state
699
+
700
+
701
+ class ConvUpscale(Conv):
702
+ """
703
+ Constructs an upscaler using a convolutional layer and pixel shuffling.
704
+
705
+ Supports 1D, 2D, and 3D convolutional upscaling.
706
+
707
+ See `Real-Time Single Image and Video Super-Resolution Using an Efficient Sub-Pixel
708
+ Convolutional Neural Network <https://arxiv.org/abs/1609.05158>`_ by Shi et al. (2016) for
709
+ details.
710
+
711
+ Attributes
712
+ ----------
713
+ group : int
714
+ Layer group, if 0 it will always be used, else it will only be used if its group matches the
715
+ Networks
716
+ description : str
717
+ Description of the layer
718
+ layers : Sequential
719
+ Layers to loop through in the forward pass
720
+ """
721
+ def __init__(
722
+ self,
723
+ net_out: list[int],
724
+ shapes: Shapes,
725
+ *,
726
+ filters: int | None = None,
727
+ layer: int | None = None,
728
+ factor: float | None = None,
729
+ scale: int = 2,
730
+ kernel: int | list[int] = 3,
731
+ dropout: float = 0,
732
+ activation: str | None = 'ELU',
733
+ norm: Literal[None, 'batch', 'layer'] = None,
734
+ padding_mode: Literal[None, 'zeros', 'reflect', 'replicate', 'circular'] = None,
735
+ **kwargs: Any) -> None:
736
+ """
737
+ Parameters
738
+ ----------
739
+ net_out : list[int]
740
+ Shape of the network's output, required only if layer contains factor
741
+ shapes: Shapes
742
+ Shape of the outputs from each layer
743
+ filters : int | None, Optional
744
+ Number of convolutional filters, will be used if provided, else factor will be used
745
+ layer : int | None, Optional
746
+ If factor is not None, which layer for factor to be relative to, if None, network output
747
+ will be used
748
+ factor : float | None, Optional
749
+ Number of convolutional filters equal to the output channels multiplied by factor,
750
+ won't be used if filters is provided
751
+ scale : int, Optional
752
+ Factor to upscale the input by, default = 2
753
+ kernel : int | list[int], Optional
754
+ Size of the kernel, default = 3
755
+ dropout : float, Optional
756
+ Probability of dropout, default = 0
757
+ activation : str | None, Optional
758
+ Which activation function to use, default = 'ELU'
759
+ norm : {None, 'batch', 'layer'}
760
+ If batch or layer normalisation should be used
761
+ padding_mode : {None, 'zeros', 'reflect', 'replicate', 'circular'}
762
+ Padding mode to use from PyTorch, if None zero padding is used if version is >3.9.4 else
763
+ replicate padding is used
764
+ **kwargs
765
+ Leftover parameters to pass to base layer for checking
766
+ """
767
+ filters_scale: int = scale ** (len(shapes[-1]) - 1)
768
+ filters = self._check_factor_filters(
769
+ shapes[-1].copy(),
770
+ filters=filters,
771
+ factor=factor,
772
+ target=shapes[layer] if layer is not None else net_out,
773
+ )[0] * filters_scale
774
+
775
+ # Convolutional layer
776
+ super().__init__(
777
+ net_out=net_out,
778
+ shapes=shapes,
779
+ filters=filters,
780
+ kernel=kernel,
781
+ stride=1,
782
+ padding='same',
783
+ dropout=dropout,
784
+ padding_mode=padding_mode,
785
+ activation=activation,
786
+ norm=norm,
787
+ **kwargs,
788
+ )
789
+
790
+ # Upscaling done using pixel shuffling
791
+ self.layers.add_module('PixelShuffle', PixelShuffle(scale))
792
+ shapes[-1][0] = shapes[-1][0] // filters_scale
793
+ shapes[-1][1:] = [length * scale for length in shapes[-1][1:]]
794
+
795
+ def __getstate__(self) -> dict[str, Any]:
796
+ state: dict[str, Any] = super().__getstate__()
797
+ state.pop('stride', None)
798
+ state.pop('padding', None)
799
+ return state
800
+
801
+
802
+ class PixelShuffle(BaseLayer):
803
+ r"""
804
+ Used for upscaling by scale factor :math:`r` for an input :math:`(N,C\times r^n,D_1,...,D_n)` to
805
+ an output :math:`(N,C,D_1\times r,...,D_n\times r)`.
806
+
807
+ Equivalent to :class:`torch.nn.PixelShuffle`, but for nD.
808
+
809
+ Attributes
810
+ ----------
811
+ group : int
812
+ Layer group, if 0 it will always be used, else it will only be used if its group matches the
813
+ Networks
814
+ description : str
815
+ Description of the layer
816
+ """
817
+ def __init__(self, scale: int, *, shapes: Shapes | None = None, **kwargs: Any) -> None:
818
+ """
819
+ Parameters
820
+ ----------
821
+ scale : int
822
+ Upscaling factor
823
+ shapes: Shapes | None, Optional
824
+ Shape of the outputs from each layer
825
+ **kwargs
826
+ Leftover parameters to pass to base layer for checking
827
+ """
828
+ super().__init__(**({'idx': 0} | kwargs))
829
+ self._scale: int = scale
830
+ filters_scale: int
831
+
832
+ # If not used as an individual layer in Network
833
+ if not shapes:
834
+ return
835
+
836
+ filters_scale = self._scale ** (len(shapes[-1][1:]))
837
+ self._check_filters(filters_scale, shapes[-1])
838
+
839
+ shapes.append(shapes[-1].copy())
840
+ shapes[-1][0] = shapes[-1][0] // filters_scale
841
+ shapes[-1][1:] = [length * self._scale for length in shapes[-1][1:]]
842
+
843
+ def __getstate__(self) -> dict[str, Any]:
844
+ return super().__getstate__() | {'scale': self._scale}
845
+
846
+ @staticmethod
847
+ def _check_filters(filters_scale: int, shape: list[int]) -> None:
848
+ """
849
+ Checks if the number of channels is an integer multiple of the upscaling factor
850
+
851
+ Parameters
852
+ ----------
853
+ filters_scale : int
854
+ Upscaling factor for the number of filters
855
+ shape : list[int]
856
+ Shape of the input
857
+ """
858
+ if shape[0] % filters_scale != 0:
859
+ raise ValueError(f'Channels ({shape}) must be an integer multiple of '
860
+ f'{filters_scale}')
861
+
862
+ def forward(self, x: Tensor, *_: Any, **__: Any) -> Tensor:
863
+ r"""
864
+ Forward pass of pixel shuffle
865
+
866
+ Parameters
867
+ ----------
868
+ x : Tensor
869
+ Input tensor with shape :math:`(N,C\times r^n,D_1,...,D_n)` and type float, where N is
870
+ the batch size, C is the number of channels, r is the upscaling factor and :math:`D_n`
871
+ is the length of dimension n
872
+
873
+ Returns
874
+ -------
875
+ Tensor
876
+ Output tensor with shape :math:`(N,C,D_1\times r,...,D_n\times r)` and type float
877
+ """
878
+ dims: int
879
+ filters_scale: int = self._scale ** (len(x.shape[2:]))
880
+ output_channels: int = x.size(1) // filters_scale
881
+ output_shape: Tensor = self._scale * torch.tensor(x.shape[2:])
882
+ idxs: Tensor
883
+
884
+ dims = len(output_shape)
885
+ idxs = torch.arange(dims * 2) + 2
886
+
887
+ x = x.view([x.size(0), output_channels, *[self._scale] * len(x.shape[2:]), *x.shape[2:]])
888
+ x = x.permute(0, 1, *torch.ravel(torch.column_stack((idxs[dims:], idxs[:dims]))))
889
+ x = x.reshape(x.size(0), output_channels, *output_shape)
890
+ return x
891
+
892
+ def extra_repr(self) -> str:
893
+ """
894
+ Displays layer parameters when printing the network
895
+
896
+ Returns
897
+ -------
898
+ str
899
+ Layer parameters
900
+ """
901
+ return f'upscale_factor={self._scale}'
902
+
903
+
904
+ def _kernel_transpose_shape(
905
+ kernel: int | list[int],
906
+ strides: int | list[int],
907
+ padding: int | list[int],
908
+ dilation: int | list[int],
909
+ out_padding: int | list[int],
910
+ shape: list[int]) -> list[int]:
911
+ """
912
+ Calculates the output shape after a transposed convolutional kernel
913
+
914
+ Parameters
915
+ ----------
916
+ kernel : int | list[int]
917
+ Size of the kernel
918
+ strides : int | list[int]
919
+ Stride of the kernel
920
+ padding : int | list[int]
921
+ Input padding
922
+ dilation : int | list[int]
923
+ Spacing between kernel elements
924
+ shape : list[int]
925
+ Input shape of the layer
926
+
927
+ Returns
928
+ -------
929
+ list[int]
930
+ Output shape of the layer
931
+ """
932
+ shape = shape.copy()
933
+ strides, kernel, padding, dilation, out_padding = _int_list_conversion(
934
+ len(shape[1:]),
935
+ [strides, kernel, padding, dilation, out_padding]
936
+ )
937
+
938
+ for i, (stride, kernel_length, pad, dilation_length, out_pad, length) in enumerate(zip(
939
+ strides,
940
+ kernel,
941
+ padding,
942
+ dilation,
943
+ out_padding,
944
+ shape[1:]
945
+ )):
946
+ shape[i + 1] = max(
947
+ 1,
948
+ stride * (length - 1) + dilation_length * (kernel_length - 1) - 2 * pad + out_pad + 1
949
+ )
950
+ return shape
951
+
952
+
953
+ def _padding_transpose(
954
+ kernel: int | list[int],
955
+ strides: int | list[int],
956
+ dilation: int | list[int],
957
+ in_shape: list[int],
958
+ out_shape: list[int]) -> list[int]:
959
+ """
960
+ Calculates the padding required for specific output shape.
961
+
962
+ Parameters
963
+ ----------
964
+ kernel : int | list[int]
965
+ Size of the kernel
966
+ strides : int | list[int]
967
+ Stride of the kernel
968
+ dilation : int | list[int]
969
+ Spacing between kernel elements
970
+ in_shape : list[int]
971
+ Input shape of the layer
972
+ out_shape : list[int]
973
+ Output shape of the layer
974
+
975
+ Returns
976
+ -------
977
+ list[int]
978
+ Required padding for specific output shape
979
+ """
980
+ padding: list[int] = []
981
+ strides, kernel, dilation = _int_list_conversion(
982
+ len(in_shape[1:]),
983
+ [strides, kernel, dilation],
984
+ )
985
+
986
+ for stride, kernel_length, dilation_length, in_length, out_length in zip(
987
+ strides,
988
+ kernel,
989
+ dilation,
990
+ in_shape[1:],
991
+ out_shape[1:],
992
+ ):
993
+ padding.append(
994
+ (stride * (in_length - 1) + dilation_length * (kernel_length - 1) - out_length + 1) // 2
995
+ )
996
+ return padding
997
+
998
+
999
+ __all__ = [
1000
+ 'Conv',
1001
+ 'ConvDepth',
1002
+ 'ConvDepthDownscale',
1003
+ 'ConvDownscale',
1004
+ 'ConvTranspose',
1005
+ 'ConvTransposeUpscale',
1006
+ 'ConvUpscale',
1007
+ 'PixelShuffle',
1008
+ ]