learnergy 2.0.0__tar.gz → 2.0.1__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.
- {learnergy-2.0.0 → learnergy-2.0.1}/PKG-INFO +30 -4
- {learnergy-2.0.0 → learnergy-2.0.1}/README.md +29 -3
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/__init__.py +1 -1
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/math/metrics.py +13 -4
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/bernoulli/discriminative_rbm.py +1 -2
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/bernoulli/dropout_rbm.py +4 -28
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/bernoulli/rbm.py +1 -6
- learnergy-2.0.1/learnergy/models/gaussian/_normalization.py +13 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/gaussian/gaussian_conv_rbm.py +4 -7
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/gaussian/gaussian_rbm.py +18 -31
- learnergy-2.0.1/learnergy/visual/tensor.py +38 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy.egg-info/PKG-INFO +30 -4
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy.egg-info/SOURCES.txt +1 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/pyproject.toml +1 -1
- {learnergy-2.0.0 → learnergy-2.0.1}/tests/test_bernoulli.py +61 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/tests/test_core.py +1 -1
- {learnergy-2.0.0 → learnergy-2.0.1}/tests/test_deep.py +48 -3
- learnergy-2.0.1/tests/test_gaussian.py +226 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/tests/test_utilities.py +52 -0
- learnergy-2.0.0/learnergy/visual/tensor.py +0 -32
- learnergy-2.0.0/tests/test_gaussian.py +0 -73
- {learnergy-2.0.0 → learnergy-2.0.1}/LICENSE +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/core/__init__.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/core/dataset.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/core/model.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/math/__init__.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/math/scale.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/__init__.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/bernoulli/__init__.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/bernoulli/conv_rbm.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/bernoulli/e_dropout_rbm.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/deep/__init__.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/deep/conv_dbn.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/deep/dbn.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/deep/residual_dbn.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/extra/__init__.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/extra/sigmoid_rbm.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/gaussian/__init__.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/utils/__init__.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/utils/constants.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/utils/exception.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/utils/logging.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/visual/__init__.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/visual/convergence.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/visual/image.py +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy.egg-info/dependency_links.txt +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy.egg-info/requires.txt +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/learnergy.egg-info/top_level.txt +0 -0
- {learnergy-2.0.0 → learnergy-2.0.1}/setup.cfg +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: learnergy
|
|
3
|
-
Version: 2.0.
|
|
3
|
+
Version: 2.0.1
|
|
4
4
|
Summary: Energy-based machine learners built with PyTorch
|
|
5
5
|
Author-email: Mateus Roder <mateus.roder@unesp.br>, Gustavo de Rosa <gustavo.rosa@unesp.br>
|
|
6
6
|
License-Expression: Apache-2.0
|
|
@@ -49,15 +49,22 @@ image-quality metrics, and visualization helpers.
|
|
|
49
49
|
|
|
50
50
|
## Installation
|
|
51
51
|
|
|
52
|
-
Learnergy requires Python 3.11 or newer.
|
|
52
|
+
Learnergy requires Python 3.11 or newer. Add it to a project managed by uv with:
|
|
53
53
|
|
|
54
54
|
```bash
|
|
55
|
-
|
|
55
|
+
uv add learnergy
|
|
56
|
+
```
|
|
57
|
+
|
|
58
|
+
Add the optional torchvision dependency to run the examples:
|
|
59
|
+
|
|
60
|
+
```bash
|
|
61
|
+
uv add "learnergy[examples]"
|
|
56
62
|
```
|
|
57
63
|
|
|
58
|
-
|
|
64
|
+
For a consumer installation in an existing Python environment, pip is also supported:
|
|
59
65
|
|
|
60
66
|
```bash
|
|
67
|
+
pip install learnergy
|
|
61
68
|
pip install "learnergy[examples]"
|
|
62
69
|
```
|
|
63
70
|
|
|
@@ -112,6 +119,25 @@ plots, image mosaics, and tensor rendering.
|
|
|
112
119
|
See [`examples/applications`](examples/applications) for complete training and
|
|
113
120
|
classification programs.
|
|
114
121
|
|
|
122
|
+
### Numerical behavior
|
|
123
|
+
|
|
124
|
+
When enabled, Gaussian normalization uses statistics from the current batch,
|
|
125
|
+
not stored training statistics. Batches of two or more samples use sample
|
|
126
|
+
standard deviation; a singleton batch is centered to zero. Representations
|
|
127
|
+
therefore depend on batch composition. Disable the corresponding normalization
|
|
128
|
+
flags when supplying externally standardized features.
|
|
129
|
+
|
|
130
|
+
`VarianceGaussianRBM.sigma` is a learnable scale: the effective visible variance
|
|
131
|
+
is `sigma**2` plus a dtype-dependent epsilon. Its `visible_sampling` method
|
|
132
|
+
returns conditional means followed by sampled states, and Gibbs sampling uses
|
|
133
|
+
those states.
|
|
134
|
+
|
|
135
|
+
Gaussian convolutional representations support gradient-based fine-tuning.
|
|
136
|
+
Use `torch.no_grad()` when extracting frozen features without an autograd graph.
|
|
137
|
+
|
|
138
|
+
The corrected variance-Gaussian sampling and stabilized likelihood calculations
|
|
139
|
+
can change training trajectories, including with a fixed random seed.
|
|
140
|
+
|
|
115
141
|
## Development
|
|
116
142
|
|
|
117
143
|
The repository uses [uv](https://docs.astral.sh/uv/) for reproducible
|
|
@@ -12,15 +12,22 @@ image-quality metrics, and visualization helpers.
|
|
|
12
12
|
|
|
13
13
|
## Installation
|
|
14
14
|
|
|
15
|
-
Learnergy requires Python 3.11 or newer.
|
|
15
|
+
Learnergy requires Python 3.11 or newer. Add it to a project managed by uv with:
|
|
16
16
|
|
|
17
17
|
```bash
|
|
18
|
-
|
|
18
|
+
uv add learnergy
|
|
19
|
+
```
|
|
20
|
+
|
|
21
|
+
Add the optional torchvision dependency to run the examples:
|
|
22
|
+
|
|
23
|
+
```bash
|
|
24
|
+
uv add "learnergy[examples]"
|
|
19
25
|
```
|
|
20
26
|
|
|
21
|
-
|
|
27
|
+
For a consumer installation in an existing Python environment, pip is also supported:
|
|
22
28
|
|
|
23
29
|
```bash
|
|
30
|
+
pip install learnergy
|
|
24
31
|
pip install "learnergy[examples]"
|
|
25
32
|
```
|
|
26
33
|
|
|
@@ -75,6 +82,25 @@ plots, image mosaics, and tensor rendering.
|
|
|
75
82
|
See [`examples/applications`](examples/applications) for complete training and
|
|
76
83
|
classification programs.
|
|
77
84
|
|
|
85
|
+
### Numerical behavior
|
|
86
|
+
|
|
87
|
+
When enabled, Gaussian normalization uses statistics from the current batch,
|
|
88
|
+
not stored training statistics. Batches of two or more samples use sample
|
|
89
|
+
standard deviation; a singleton batch is centered to zero. Representations
|
|
90
|
+
therefore depend on batch composition. Disable the corresponding normalization
|
|
91
|
+
flags when supplying externally standardized features.
|
|
92
|
+
|
|
93
|
+
`VarianceGaussianRBM.sigma` is a learnable scale: the effective visible variance
|
|
94
|
+
is `sigma**2` plus a dtype-dependent epsilon. Its `visible_sampling` method
|
|
95
|
+
returns conditional means followed by sampled states, and Gibbs sampling uses
|
|
96
|
+
those states.
|
|
97
|
+
|
|
98
|
+
Gaussian convolutional representations support gradient-based fine-tuning.
|
|
99
|
+
Use `torch.no_grad()` when extracting frozen features without an autograd graph.
|
|
100
|
+
|
|
101
|
+
The corrected variance-Gaussian sampling and stabilized likelihood calculations
|
|
102
|
+
can change training trajectories, including with a fixed random seed.
|
|
103
|
+
|
|
78
104
|
## Development
|
|
79
105
|
|
|
80
106
|
The repository uses [uv](https://docs.astral.sh/uv/) for reproducible
|
|
@@ -5,17 +5,26 @@ from skimage.metrics import structural_similarity
|
|
|
5
5
|
|
|
6
6
|
|
|
7
7
|
def calculate_ssim(v: torch.Tensor, x: torch.Tensor) -> float:
|
|
8
|
-
"""Calculate the mean structural similarity of reconstructed images.
|
|
8
|
+
"""Calculate the mean structural similarity of reconstructed images.
|
|
9
|
+
|
|
10
|
+
Args:
|
|
11
|
+
v: Reconstructed images, with each image flattened or shaped like its original.
|
|
12
|
+
x: Original grayscale images with shape (batch, height, width).
|
|
13
|
+
|
|
14
|
+
Raises:
|
|
15
|
+
ValueError: The batches contain different numbers of images.
|
|
16
|
+
|
|
17
|
+
"""
|
|
9
18
|
|
|
10
19
|
originals = x.detach().cpu().numpy()
|
|
11
20
|
reconstructed = v.detach().cpu().numpy()
|
|
12
|
-
|
|
21
|
+
height, width = originals.shape[1:3]
|
|
13
22
|
|
|
14
23
|
return sum(
|
|
15
24
|
structural_similarity(
|
|
16
25
|
original,
|
|
17
|
-
rebuilt.reshape(
|
|
26
|
+
rebuilt.reshape(height, width),
|
|
18
27
|
data_range=original.max() - original.min(),
|
|
19
28
|
)
|
|
20
|
-
for original, rebuilt in zip(originals, reconstructed)
|
|
29
|
+
for original, rebuilt in zip(originals, reconstructed, strict=True)
|
|
21
30
|
) / len(reconstructed)
|
|
@@ -288,8 +288,7 @@ class HybridDiscriminativeRBM(DiscriminativeRBM):
|
|
|
288
288
|
|
|
289
289
|
"""
|
|
290
290
|
|
|
291
|
-
|
|
292
|
-
probs = torch.div(activations, torch.sum(activations, dim=1).unsqueeze(1))
|
|
291
|
+
probs = F.softmax(F.linear(h, self.U, self.c), dim=1)
|
|
293
292
|
states = torch.nn.functional.one_hot(
|
|
294
293
|
torch.argmax(probs, dim=1), num_classes=self.n_classes
|
|
295
294
|
).float()
|
|
@@ -4,7 +4,6 @@ from typing import Tuple
|
|
|
4
4
|
|
|
5
5
|
import torch
|
|
6
6
|
import torch.nn.functional as F
|
|
7
|
-
from torch.utils.data import DataLoader
|
|
8
7
|
|
|
9
8
|
import learnergy.utils.exception as e
|
|
10
9
|
from learnergy.core.model import _validated_property
|
|
@@ -115,35 +114,12 @@ class DropoutRBM(RBM):
|
|
|
115
114
|
|
|
116
115
|
"""
|
|
117
116
|
|
|
118
|
-
mse = 0
|
|
119
|
-
batch_size = len(dataset)
|
|
120
|
-
|
|
121
|
-
# Saving dropout rate to an auxiliary variable
|
|
122
|
-
# and temporarily disabling dropout
|
|
123
117
|
p = self.p
|
|
124
118
|
self.p = 0
|
|
125
|
-
|
|
126
|
-
|
|
127
|
-
|
|
128
|
-
|
|
129
|
-
|
|
130
|
-
for samples, _ in batches:
|
|
131
|
-
samples = samples.reshape(len(samples), self.n_visible).to(self.device)
|
|
132
|
-
|
|
133
|
-
_, pos_hidden_states = self.hidden_sampling(samples)
|
|
134
|
-
visible_probs, visible_states = self.visible_sampling(pos_hidden_states)
|
|
135
|
-
|
|
136
|
-
batch_mse = torch.div(
|
|
137
|
-
torch.sum(torch.pow(samples - visible_states, 2)), batch_size
|
|
138
|
-
)
|
|
139
|
-
mse += batch_mse
|
|
140
|
-
|
|
141
|
-
mse /= len(batches)
|
|
142
|
-
|
|
143
|
-
# Recovering initial dropout rate
|
|
144
|
-
self.p = p
|
|
145
|
-
|
|
146
|
-
return mse, visible_probs
|
|
119
|
+
try:
|
|
120
|
+
return super().reconstruct(dataset)
|
|
121
|
+
finally:
|
|
122
|
+
self.p = p
|
|
147
123
|
|
|
148
124
|
|
|
149
125
|
class DropConnectRBM(DropoutRBM):
|
|
@@ -247,12 +247,7 @@ class RBM(Model):
|
|
|
247
247
|
energy1 = self.energy(samples_binary)
|
|
248
248
|
|
|
249
249
|
# Calculate the logarithm of the pseudo-likelihood
|
|
250
|
-
pl = torch.mean(
|
|
251
|
-
self.n_visible
|
|
252
|
-
* torch.log(
|
|
253
|
-
torch.sigmoid(energy1 - energy) + torch.finfo(samples.dtype).eps
|
|
254
|
-
)
|
|
255
|
-
)
|
|
250
|
+
pl = torch.mean(self.n_visible * F.logsigmoid(energy1 - energy))
|
|
256
251
|
|
|
257
252
|
return pl
|
|
258
253
|
|
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
"""Batch standardization shared by Gaussian models."""
|
|
2
|
+
|
|
3
|
+
import torch
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
def standardize(samples: torch.Tensor) -> torch.Tensor:
|
|
7
|
+
"""Use sample variance for full batches and center singleton batches at zero."""
|
|
8
|
+
|
|
9
|
+
correction = 1 if len(samples) > 1 else 0
|
|
10
|
+
std = samples.std(dim=0, correction=correction, keepdim=True)
|
|
11
|
+
return (samples - samples.mean(dim=0, keepdim=True)) / (
|
|
12
|
+
std + torch.finfo(samples.dtype).eps
|
|
13
|
+
)
|
|
@@ -8,6 +8,7 @@ from torch.utils.data import DataLoader
|
|
|
8
8
|
|
|
9
9
|
from learnergy.core.model import _validated_property
|
|
10
10
|
from learnergy.models.bernoulli.conv_rbm import ConvRBM
|
|
11
|
+
from learnergy.models.gaussian._normalization import standardize
|
|
11
12
|
|
|
12
13
|
|
|
13
14
|
class GaussianConvRBM(ConvRBM):
|
|
@@ -51,7 +52,7 @@ class GaussianConvRBM(ConvRBM):
|
|
|
51
52
|
"""Compute hidden probabilities and activations."""
|
|
52
53
|
|
|
53
54
|
activations = F.conv2d(v, self.W, bias=self.b)
|
|
54
|
-
return F.relu6(activations)
|
|
55
|
+
return F.relu6(activations), activations
|
|
55
56
|
|
|
56
57
|
def visible_sampling(self, h: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
|
57
58
|
"""Compute visible probabilities and activations."""
|
|
@@ -86,10 +87,7 @@ class GaussianConvRBM(ConvRBM):
|
|
|
86
87
|
).to(self.device)
|
|
87
88
|
|
|
88
89
|
if self.normalize:
|
|
89
|
-
|
|
90
|
-
samples = (samples - samples.mean(0, True)) / (
|
|
91
|
-
samples.std(0, True) + eps
|
|
92
|
-
)
|
|
90
|
+
samples = standardize(samples)
|
|
93
91
|
|
|
94
92
|
_, _, _, _, visible_states = self.gibbs_sampling(samples)
|
|
95
93
|
visible_states = visible_states.detach()
|
|
@@ -110,8 +108,7 @@ class GaussianConvRBM(ConvRBM):
|
|
|
110
108
|
"""Return hidden activations, optionally pooled."""
|
|
111
109
|
|
|
112
110
|
if self.normalize:
|
|
113
|
-
|
|
114
|
-
x = (x - x.mean(0, True)) / (x.std(0, True) + eps)
|
|
111
|
+
x = standardize(x)
|
|
115
112
|
|
|
116
113
|
x, _ = self.hidden_sampling(x)
|
|
117
114
|
if self.maxpooling:
|
|
@@ -10,6 +10,7 @@ from torch.utils.data import DataLoader
|
|
|
10
10
|
|
|
11
11
|
from learnergy.core.model import _validated_property
|
|
12
12
|
from learnergy.models.bernoulli.rbm import RBM
|
|
13
|
+
from learnergy.models.gaussian._normalization import standardize
|
|
13
14
|
|
|
14
15
|
|
|
15
16
|
class GaussianRBM(RBM):
|
|
@@ -150,10 +151,7 @@ class GaussianRBM(RBM):
|
|
|
150
151
|
|
|
151
152
|
for samples, _ in batches:
|
|
152
153
|
if self.normalize:
|
|
153
|
-
samples = (
|
|
154
|
-
(samples - torch.mean(samples, 0, True))
|
|
155
|
-
/ (torch.std(samples, 0, True) + torch.finfo(samples.dtype).eps)
|
|
156
|
-
).detach()
|
|
154
|
+
samples = standardize(samples).detach()
|
|
157
155
|
|
|
158
156
|
samples = samples.reshape(len(samples), self.n_visible).to(self.device)
|
|
159
157
|
|
|
@@ -210,10 +208,7 @@ class GaussianRBM(RBM):
|
|
|
210
208
|
|
|
211
209
|
for samples, _ in batches:
|
|
212
210
|
if self.normalize:
|
|
213
|
-
samples = (
|
|
214
|
-
(samples - torch.mean(samples, 0, True))
|
|
215
|
-
/ (torch.std(samples, 0, True) + torch.finfo(samples.dtype).eps)
|
|
216
|
-
).detach()
|
|
211
|
+
samples = standardize(samples).detach()
|
|
217
212
|
|
|
218
213
|
samples = samples.reshape(len(samples), self.n_visible).to(self.device)
|
|
219
214
|
|
|
@@ -241,10 +236,7 @@ class GaussianRBM(RBM):
|
|
|
241
236
|
"""
|
|
242
237
|
|
|
243
238
|
if self.input_normalize:
|
|
244
|
-
x = (
|
|
245
|
-
(x - torch.mean(x, 0, True))
|
|
246
|
-
/ (torch.std(x, 0, True) + torch.finfo(x.dtype).eps)
|
|
247
|
-
).detach()
|
|
239
|
+
x = standardize(x).detach()
|
|
248
240
|
|
|
249
241
|
x, _ = self.hidden_sampling(x)
|
|
250
242
|
|
|
@@ -421,8 +413,9 @@ class VarianceGaussianRBM(RBM):
|
|
|
421
413
|
"""A VarianceGaussianRBM class provides the basic implementation for
|
|
422
414
|
Gaussian-Bernoulli Restricted Boltzmann Machines (without standardization).
|
|
423
415
|
|
|
424
|
-
|
|
425
|
-
|
|
416
|
+
The learnable scale parameter ``sigma`` defines the visible variance as
|
|
417
|
+
``sigma**2 + torch.finfo(dtype).eps``. The same variance is used by the
|
|
418
|
+
free energy and the visible conditional distribution.
|
|
426
419
|
|
|
427
420
|
Therefore, there is no need to standardize the data, as the variance
|
|
428
421
|
will be trained throughout the learning procedure.
|
|
@@ -490,8 +483,8 @@ class VarianceGaussianRBM(RBM):
|
|
|
490
483
|
|
|
491
484
|
"""
|
|
492
485
|
|
|
493
|
-
|
|
494
|
-
activations = F.linear(
|
|
486
|
+
variance = self.sigma.square() + torch.finfo(v.dtype).eps
|
|
487
|
+
activations = F.linear(v / variance, self.W.t(), self.b)
|
|
495
488
|
|
|
496
489
|
if scale:
|
|
497
490
|
probs = torch.sigmoid(torch.div(activations, self.T))
|
|
@@ -512,22 +505,16 @@ class VarianceGaussianRBM(RBM):
|
|
|
512
505
|
scale: A boolean to decide whether temperature should be used or not.
|
|
513
506
|
|
|
514
507
|
Returns:
|
|
515
|
-
The
|
|
508
|
+
The conditional means and sampled visible states, respectively.
|
|
516
509
|
|
|
517
510
|
"""
|
|
518
511
|
|
|
519
512
|
activations = F.linear(h, self.W, self.a)
|
|
513
|
+
variance = self.sigma.square() + torch.finfo(activations.dtype).eps
|
|
514
|
+
std = variance.sqrt().expand_as(activations)
|
|
515
|
+
states = torch.normal(activations, std)
|
|
520
516
|
|
|
521
|
-
|
|
522
|
-
# Variance needs to have size equal to (batch_size, n_visible)
|
|
523
|
-
sigma = self.sigma.unsqueeze(0).expand(activations.size(0), -1)
|
|
524
|
-
else:
|
|
525
|
-
# Variance needs to have size equal to (n_visible)
|
|
526
|
-
sigma = self.sigma
|
|
527
|
-
|
|
528
|
-
states = torch.normal(activations, torch.pow(sigma, 2))
|
|
529
|
-
|
|
530
|
-
return states, activations
|
|
517
|
+
return activations, states
|
|
531
518
|
|
|
532
519
|
def energy(self, samples: torch.Tensor) -> torch.Tensor:
|
|
533
520
|
"""Calculates and frees the system's energy.
|
|
@@ -540,13 +527,13 @@ class VarianceGaussianRBM(RBM):
|
|
|
540
527
|
|
|
541
528
|
"""
|
|
542
529
|
|
|
543
|
-
|
|
544
|
-
activations = F.linear(
|
|
530
|
+
variance = self.sigma.square() + torch.finfo(samples.dtype).eps
|
|
531
|
+
activations = F.linear(samples / variance, self.W.t(), self.b)
|
|
545
532
|
|
|
546
533
|
h = torch.sum(F.softplus(activations), dim=1)
|
|
547
|
-
v = torch.sum(
|
|
534
|
+
v = torch.sum((samples - self.a).square() / (2 * variance), dim=1)
|
|
548
535
|
|
|
549
|
-
energy =
|
|
536
|
+
energy = v - h
|
|
550
537
|
|
|
551
538
|
return energy
|
|
552
539
|
|
|
@@ -0,0 +1,38 @@
|
|
|
1
|
+
"""Tensor visualization."""
|
|
2
|
+
|
|
3
|
+
import matplotlib.pyplot as plt
|
|
4
|
+
import torch
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def _show(tensor: torch.Tensor) -> None:
|
|
8
|
+
image = tensor
|
|
9
|
+
if tensor.ndim == 3:
|
|
10
|
+
image = tensor.permute(1, 2, 0) if tensor.size(0) == 3 else tensor.squeeze(0)
|
|
11
|
+
plt.imshow(
|
|
12
|
+
image.detach().cpu().numpy(),
|
|
13
|
+
cmap=None if image.ndim == 3 else "gray",
|
|
14
|
+
)
|
|
15
|
+
plt.xticks([])
|
|
16
|
+
plt.yticks([])
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def save_tensor(tensor: torch.Tensor, output_path: str) -> None:
|
|
20
|
+
"""Save an (H, W), (1, H, W), or (3, H, W) image tensor."""
|
|
21
|
+
|
|
22
|
+
figure = plt.figure()
|
|
23
|
+
try:
|
|
24
|
+
_show(tensor)
|
|
25
|
+
figure.savefig(output_path)
|
|
26
|
+
finally:
|
|
27
|
+
plt.close(figure)
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def show_tensor(tensor: torch.Tensor) -> None:
|
|
31
|
+
"""Display an (H, W), (1, H, W), or (3, H, W) image tensor."""
|
|
32
|
+
|
|
33
|
+
figure = plt.figure()
|
|
34
|
+
try:
|
|
35
|
+
_show(tensor)
|
|
36
|
+
plt.show()
|
|
37
|
+
finally:
|
|
38
|
+
plt.close(figure)
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: learnergy
|
|
3
|
-
Version: 2.0.
|
|
3
|
+
Version: 2.0.1
|
|
4
4
|
Summary: Energy-based machine learners built with PyTorch
|
|
5
5
|
Author-email: Mateus Roder <mateus.roder@unesp.br>, Gustavo de Rosa <gustavo.rosa@unesp.br>
|
|
6
6
|
License-Expression: Apache-2.0
|
|
@@ -49,15 +49,22 @@ image-quality metrics, and visualization helpers.
|
|
|
49
49
|
|
|
50
50
|
## Installation
|
|
51
51
|
|
|
52
|
-
Learnergy requires Python 3.11 or newer.
|
|
52
|
+
Learnergy requires Python 3.11 or newer. Add it to a project managed by uv with:
|
|
53
53
|
|
|
54
54
|
```bash
|
|
55
|
-
|
|
55
|
+
uv add learnergy
|
|
56
|
+
```
|
|
57
|
+
|
|
58
|
+
Add the optional torchvision dependency to run the examples:
|
|
59
|
+
|
|
60
|
+
```bash
|
|
61
|
+
uv add "learnergy[examples]"
|
|
56
62
|
```
|
|
57
63
|
|
|
58
|
-
|
|
64
|
+
For a consumer installation in an existing Python environment, pip is also supported:
|
|
59
65
|
|
|
60
66
|
```bash
|
|
67
|
+
pip install learnergy
|
|
61
68
|
pip install "learnergy[examples]"
|
|
62
69
|
```
|
|
63
70
|
|
|
@@ -112,6 +119,25 @@ plots, image mosaics, and tensor rendering.
|
|
|
112
119
|
See [`examples/applications`](examples/applications) for complete training and
|
|
113
120
|
classification programs.
|
|
114
121
|
|
|
122
|
+
### Numerical behavior
|
|
123
|
+
|
|
124
|
+
When enabled, Gaussian normalization uses statistics from the current batch,
|
|
125
|
+
not stored training statistics. Batches of two or more samples use sample
|
|
126
|
+
standard deviation; a singleton batch is centered to zero. Representations
|
|
127
|
+
therefore depend on batch composition. Disable the corresponding normalization
|
|
128
|
+
flags when supplying externally standardized features.
|
|
129
|
+
|
|
130
|
+
`VarianceGaussianRBM.sigma` is a learnable scale: the effective visible variance
|
|
131
|
+
is `sigma**2` plus a dtype-dependent epsilon. Its `visible_sampling` method
|
|
132
|
+
returns conditional means followed by sampled states, and Gibbs sampling uses
|
|
133
|
+
those states.
|
|
134
|
+
|
|
135
|
+
Gaussian convolutional representations support gradient-based fine-tuning.
|
|
136
|
+
Use `torch.no_grad()` when extracting frozen features without an autograd graph.
|
|
137
|
+
|
|
138
|
+
The corrected variance-Gaussian sampling and stabilized likelihood calculations
|
|
139
|
+
can change training trajectories, including with a fixed random seed.
|
|
140
|
+
|
|
115
141
|
## Development
|
|
116
142
|
|
|
117
143
|
The repository uses [uv](https://docs.astral.sh/uv/) for reproducible
|
|
@@ -27,6 +27,7 @@ learnergy/models/deep/residual_dbn.py
|
|
|
27
27
|
learnergy/models/extra/__init__.py
|
|
28
28
|
learnergy/models/extra/sigmoid_rbm.py
|
|
29
29
|
learnergy/models/gaussian/__init__.py
|
|
30
|
+
learnergy/models/gaussian/_normalization.py
|
|
30
31
|
learnergy/models/gaussian/gaussian_conv_rbm.py
|
|
31
32
|
learnergy/models/gaussian/gaussian_rbm.py
|
|
32
33
|
learnergy/utils/__init__.py
|
|
@@ -1,3 +1,5 @@
|
|
|
1
|
+
import math
|
|
2
|
+
|
|
1
3
|
import pytest
|
|
2
4
|
import torch
|
|
3
5
|
from torch.utils.data import TensorDataset
|
|
@@ -171,3 +173,62 @@ def test_model_initialization_restores_float32_default():
|
|
|
171
173
|
assert model.W.dtype == torch.float32
|
|
172
174
|
finally:
|
|
173
175
|
torch.set_default_dtype(torch.float32)
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
@pytest.mark.parametrize("bias", [-100.0, 0.0, 100.0])
|
|
179
|
+
def test_rbm_pseudo_likelihood_uses_stable_log_probabilities(bias):
|
|
180
|
+
model = RBM(n_visible=1, n_hidden=1)
|
|
181
|
+
with torch.no_grad():
|
|
182
|
+
model.W.zero_()
|
|
183
|
+
model.a.fill_(bias)
|
|
184
|
+
|
|
185
|
+
log_probability = model.pseudo_likelihood(torch.zeros(1, 1))
|
|
186
|
+
|
|
187
|
+
assert log_probability.item() == pytest.approx(-math.log1p(math.exp(bias)))
|
|
188
|
+
assert log_probability <= 0
|
|
189
|
+
log_probability.backward()
|
|
190
|
+
assert model.a.grad.item() == pytest.approx(-1 / (1 + math.exp(-bias)))
|
|
191
|
+
|
|
192
|
+
|
|
193
|
+
@pytest.mark.parametrize("offset", [-1000.0, 0.0, 1000.0])
|
|
194
|
+
def test_hybrid_class_sampling_handles_large_logits(offset):
|
|
195
|
+
model = HybridDiscriminativeRBM(n_visible=2, n_hidden=1, n_classes=3)
|
|
196
|
+
with torch.no_grad():
|
|
197
|
+
model.U.zero_()
|
|
198
|
+
model.c.copy_(torch.tensor([offset, offset + 1, offset - 1]))
|
|
199
|
+
|
|
200
|
+
probabilities, states = model.class_sampling(torch.zeros(2, 1))
|
|
201
|
+
weights = torch.tensor([1.0, math.e, 1 / math.e])
|
|
202
|
+
expected = (weights / weights.sum()).expand(2, -1)
|
|
203
|
+
|
|
204
|
+
torch.testing.assert_close(probabilities, expected)
|
|
205
|
+
torch.testing.assert_close(states, torch.tensor([[0.0, 1.0, 0.0]]).repeat(2, 1))
|
|
206
|
+
|
|
207
|
+
|
|
208
|
+
@pytest.mark.parametrize("model_class", [DropoutRBM, DropConnectRBM])
|
|
209
|
+
@pytest.mark.parametrize("shape, error", [((0, 4), ValueError), ((1, 3), RuntimeError)])
|
|
210
|
+
def test_dropout_restores_rate_after_reconstruction_failure(model_class, shape, error):
|
|
211
|
+
model = model_class(n_visible=4, n_hidden=2, dropout=0.5)
|
|
212
|
+
dataset = TensorDataset(torch.zeros(shape), torch.zeros(shape[0]))
|
|
213
|
+
|
|
214
|
+
with pytest.raises(error):
|
|
215
|
+
model.reconstruct(dataset)
|
|
216
|
+
|
|
217
|
+
assert model.p == 0.5
|
|
218
|
+
|
|
219
|
+
|
|
220
|
+
@pytest.mark.parametrize("model_class", [DropoutRBM, DropConnectRBM])
|
|
221
|
+
def test_dropout_reconstruction_matches_disabled_dropout(model_class):
|
|
222
|
+
model = model_class(n_visible=4, n_hidden=2, dropout=0.5)
|
|
223
|
+
reference = model_class(n_visible=4, n_hidden=2, dropout=0)
|
|
224
|
+
reference.load_state_dict(model.state_dict())
|
|
225
|
+
dataset = TensorDataset(torch.rand(3, 4), torch.zeros(3))
|
|
226
|
+
|
|
227
|
+
torch.manual_seed(12)
|
|
228
|
+
expected_mse, expected_reconstruction = reference.reconstruct(dataset)
|
|
229
|
+
torch.manual_seed(12)
|
|
230
|
+
mse, reconstruction = model.reconstruct(dataset)
|
|
231
|
+
|
|
232
|
+
torch.testing.assert_close(mse, expected_mse, rtol=0, atol=0)
|
|
233
|
+
torch.testing.assert_close(reconstruction, expected_reconstruction, rtol=0, atol=0)
|
|
234
|
+
assert model.p == 0.5
|
|
@@ -111,8 +111,9 @@ def test_residual_dbn_validates_weights():
|
|
|
111
111
|
ResidualDBN(zetta2=-1)
|
|
112
112
|
|
|
113
113
|
|
|
114
|
-
|
|
115
|
-
|
|
114
|
+
@pytest.mark.parametrize("n_samples", [8, 5])
|
|
115
|
+
def test_conv_dbn_builds_trains_and_reconstructs(n_samples):
|
|
116
|
+
dataset = TensorDataset(torch.rand(n_samples, 1, 8, 8), torch.zeros(n_samples))
|
|
116
117
|
model = ConvDBN(
|
|
117
118
|
visible_shape=(8, 8),
|
|
118
119
|
filter_shape=((3, 3), (3, 3)),
|
|
@@ -128,7 +129,7 @@ def test_conv_dbn_builds_trains_and_reconstructs():
|
|
|
128
129
|
|
|
129
130
|
assert len(mse) == 2
|
|
130
131
|
assert reconstruction_mse >= 0
|
|
131
|
-
assert reconstruction.shape == (
|
|
132
|
+
assert reconstruction.shape == (n_samples, 1, 8, 8)
|
|
132
133
|
assert model(dataset.tensors[0][:2]).shape == (2, 3, 4, 4)
|
|
133
134
|
|
|
134
135
|
|
|
@@ -149,6 +150,7 @@ def test_conv_dbn_pooling_configuration():
|
|
|
149
150
|
assert len(model.maxpol2d) == 2
|
|
150
151
|
samples = torch.rand(2, 1, 8, 8)
|
|
151
152
|
output = model(samples)
|
|
153
|
+
assert output.shape == (2, 3, 3, 3)
|
|
152
154
|
expected, _ = model.models[0].hidden_sampling(samples)
|
|
153
155
|
expected = model.models[0].maxpol2d(expected)
|
|
154
156
|
expected, _ = model.models[1].hidden_sampling(expected)
|
|
@@ -160,3 +162,46 @@ def test_conv_dbn_accepts_legacy_defaults():
|
|
|
160
162
|
assert model.maxpooling == (False,)
|
|
161
163
|
assert model.pooling_kernel == (2,)
|
|
162
164
|
assert signature(ConvDBN.fit).parameters["epochs"].default == (10, 10)
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
def test_gaussian_dbn_training_handles_singleton_batches():
|
|
168
|
+
model = DBN(
|
|
169
|
+
model=("gaussian", "gaussian"),
|
|
170
|
+
n_visible=4,
|
|
171
|
+
n_hidden=(3, 2),
|
|
172
|
+
steps=(1, 1),
|
|
173
|
+
learning_rate=(0.01, 0.01),
|
|
174
|
+
momentum=(0, 0),
|
|
175
|
+
decay=(0, 0),
|
|
176
|
+
temperature=(1, 1),
|
|
177
|
+
)
|
|
178
|
+
dataset = TensorDataset(torch.rand(5, 4), torch.zeros(5))
|
|
179
|
+
|
|
180
|
+
mse, pl = model.fit(dataset, batch_size=2, epochs=(1, 1))
|
|
181
|
+
|
|
182
|
+
assert all(torch.isfinite(value) for value in mse + pl)
|
|
183
|
+
assert all(torch.isfinite(parameter).all() for parameter in model.parameters())
|
|
184
|
+
|
|
185
|
+
|
|
186
|
+
def test_conv_dbn_forward_preserves_gradient_flow():
|
|
187
|
+
model = ConvDBN(
|
|
188
|
+
visible_shape=(8, 8),
|
|
189
|
+
filter_shape=((3, 3), (2, 2)),
|
|
190
|
+
n_filters=(2, 3),
|
|
191
|
+
steps=(1, 1),
|
|
192
|
+
learning_rate=(0.1, 0.1),
|
|
193
|
+
momentum=(0, 0),
|
|
194
|
+
decay=(0, 0),
|
|
195
|
+
maxpooling=(True, False),
|
|
196
|
+
)
|
|
197
|
+
with torch.no_grad():
|
|
198
|
+
for layer in model.models:
|
|
199
|
+
layer.W.fill_(0.05)
|
|
200
|
+
layer.b.fill_(0.1)
|
|
201
|
+
|
|
202
|
+
model(torch.ones(2, 1, 8, 8)).sum().backward()
|
|
203
|
+
|
|
204
|
+
for layer in model.models:
|
|
205
|
+
assert layer.W.grad is not None
|
|
206
|
+
assert torch.isfinite(layer.W.grad).all()
|
|
207
|
+
assert torch.count_nonzero(layer.W.grad) > 0
|
|
@@ -0,0 +1,226 @@
|
|
|
1
|
+
import pytest
|
|
2
|
+
import torch
|
|
3
|
+
from torch.utils.data import TensorDataset
|
|
4
|
+
|
|
5
|
+
from learnergy.models.gaussian import (
|
|
6
|
+
GaussianConvRBM,
|
|
7
|
+
GaussianConvRBM4Deep,
|
|
8
|
+
GaussianRBM,
|
|
9
|
+
GaussianRBM4deep,
|
|
10
|
+
GaussianReluRBM,
|
|
11
|
+
GaussianReluRBM4deep,
|
|
12
|
+
GaussianSeluRBM,
|
|
13
|
+
VarianceGaussianRBM,
|
|
14
|
+
)
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def test_gaussian_rbm_end_to_end():
|
|
18
|
+
torch.manual_seed(0)
|
|
19
|
+
dataset = TensorDataset(torch.rand(12, 16), torch.zeros(12))
|
|
20
|
+
model = GaussianRBM(n_visible=16, n_hidden=8)
|
|
21
|
+
|
|
22
|
+
mse, pl = model.fit(dataset, batch_size=4, epochs=1)
|
|
23
|
+
reconstruction_mse, reconstruction = model.reconstruct(dataset)
|
|
24
|
+
|
|
25
|
+
assert mse >= 0
|
|
26
|
+
assert torch.isfinite(pl)
|
|
27
|
+
assert reconstruction_mse >= 0
|
|
28
|
+
assert reconstruction.shape == (12, 16)
|
|
29
|
+
assert model(torch.rand(2, 16)).shape == (2, 8)
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
@pytest.mark.parametrize("model_class", [GaussianReluRBM, GaussianSeluRBM])
|
|
33
|
+
def test_gaussian_activation_variants(model_class):
|
|
34
|
+
model = model_class(n_visible=16, n_hidden=8)
|
|
35
|
+
probs, states = model.hidden_sampling(torch.rand(2, 16), scale=True)
|
|
36
|
+
|
|
37
|
+
assert probs.shape == states.shape == (2, 8)
|
|
38
|
+
assert torch.equal(probs, states)
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def test_deep_model_names_remain_available():
|
|
42
|
+
assert issubclass(GaussianRBM4deep, GaussianRBM)
|
|
43
|
+
assert issubclass(GaussianReluRBM4deep, GaussianReluRBM)
|
|
44
|
+
assert issubclass(GaussianConvRBM4Deep, GaussianConvRBM)
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def test_variance_gaussian_rbm_sampling_is_finite():
|
|
48
|
+
model = VarianceGaussianRBM(n_visible=16, n_hidden=8)
|
|
49
|
+
with torch.no_grad():
|
|
50
|
+
model.sigma.fill_(1e-12)
|
|
51
|
+
|
|
52
|
+
samples = torch.rand(2, 16)
|
|
53
|
+
hidden_probs, hidden_states = model.hidden_sampling(samples)
|
|
54
|
+
visible_probs, visible_states = model.visible_sampling(hidden_states)
|
|
55
|
+
|
|
56
|
+
assert torch.isfinite(hidden_probs).all()
|
|
57
|
+
assert torch.isfinite(model.energy(samples)).all()
|
|
58
|
+
assert visible_probs.shape == visible_states.shape == (2, 16)
|
|
59
|
+
assert "sigma" in model.state_dict()
|
|
60
|
+
assert len(model.optimizer.param_groups) == 2
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def test_gaussian_conv_rbm_end_to_end():
|
|
64
|
+
dataset = TensorDataset(torch.rand(8, 1, 8, 8), torch.zeros(8))
|
|
65
|
+
model = GaussianConvRBM(
|
|
66
|
+
visible_shape=(8, 8),
|
|
67
|
+
filter_shape=(3, 3),
|
|
68
|
+
n_filters=2,
|
|
69
|
+
n_channels=1,
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
assert model.fit(dataset, batch_size=4, epochs=1) >= 0
|
|
73
|
+
assert model(dataset.tensors[0][:2]).shape == (2, 2, 6, 6)
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
@pytest.mark.parametrize("model_class", [GaussianRBM, GaussianReluRBM4deep])
|
|
77
|
+
@pytest.mark.parametrize("batch_size", [1, 3])
|
|
78
|
+
def test_gaussian_forward_standardizes_batches(model_class, batch_size):
|
|
79
|
+
model = model_class(n_visible=4, n_hidden=2)
|
|
80
|
+
samples = torch.arange(batch_size * 4, dtype=torch.float32).reshape(batch_size, 4)
|
|
81
|
+
if batch_size == 1:
|
|
82
|
+
standardized = torch.zeros_like(samples)
|
|
83
|
+
else:
|
|
84
|
+
standardized = (samples - samples.mean(0, True)) / (
|
|
85
|
+
samples.std(0) + torch.finfo(samples.dtype).eps
|
|
86
|
+
)
|
|
87
|
+
expected, _ = model.hidden_sampling(standardized)
|
|
88
|
+
|
|
89
|
+
torch.testing.assert_close(model(samples), expected, rtol=0, atol=0)
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
@pytest.mark.parametrize("model_class", [GaussianRBM, GaussianReluRBM4deep])
|
|
93
|
+
def test_gaussian_handles_singleton_training_batch_and_reconstruction(model_class):
|
|
94
|
+
torch.manual_seed(0)
|
|
95
|
+
model = model_class(n_visible=4, n_hidden=2)
|
|
96
|
+
dataset = TensorDataset(torch.rand(5, 4), torch.zeros(5))
|
|
97
|
+
|
|
98
|
+
mse, pl = model.fit(dataset, batch_size=2, epochs=1)
|
|
99
|
+
reconstruction_mse, reconstruction = model.reconstruct(
|
|
100
|
+
TensorDataset(dataset.tensors[0][:1], dataset.tensors[1][:1])
|
|
101
|
+
)
|
|
102
|
+
|
|
103
|
+
assert torch.isfinite(mse) and torch.isfinite(pl)
|
|
104
|
+
assert torch.isfinite(reconstruction_mse)
|
|
105
|
+
assert torch.isfinite(reconstruction).all()
|
|
106
|
+
assert all(torch.isfinite(parameter).all() for parameter in model.parameters())
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
@pytest.mark.parametrize("model_class", [GaussianConvRBM, GaussianConvRBM4Deep])
|
|
110
|
+
@pytest.mark.parametrize("batch_size", [1, 3])
|
|
111
|
+
def test_gaussian_conv_forward_standardizes_batches(model_class, batch_size):
|
|
112
|
+
model = model_class(
|
|
113
|
+
visible_shape=(4, 4), filter_shape=(2, 2), n_filters=2, maxpooling=True
|
|
114
|
+
)
|
|
115
|
+
samples = torch.arange(batch_size * 16, dtype=torch.float32).reshape(
|
|
116
|
+
batch_size, 1, 4, 4
|
|
117
|
+
)
|
|
118
|
+
if batch_size == 1:
|
|
119
|
+
standardized = torch.zeros_like(samples)
|
|
120
|
+
else:
|
|
121
|
+
standardized = (samples - samples.mean(0, True)) / (
|
|
122
|
+
samples.std(0) + torch.finfo(samples.dtype).eps
|
|
123
|
+
)
|
|
124
|
+
expected, _ = model.hidden_sampling(standardized)
|
|
125
|
+
expected = model.maxpol2d(expected)
|
|
126
|
+
|
|
127
|
+
torch.testing.assert_close(model(samples), expected, rtol=0, atol=0)
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
@pytest.mark.parametrize("model_class", [GaussianConvRBM, GaussianConvRBM4Deep])
|
|
131
|
+
def test_gaussian_conv_handles_singleton_training_batch(model_class):
|
|
132
|
+
torch.manual_seed(0)
|
|
133
|
+
model = model_class(visible_shape=(4, 4), filter_shape=(2, 2), n_filters=2)
|
|
134
|
+
dataset = TensorDataset(torch.rand(5, 1, 4, 4), torch.zeros(5))
|
|
135
|
+
|
|
136
|
+
mse = model.fit(dataset, batch_size=2, epochs=1)
|
|
137
|
+
|
|
138
|
+
assert torch.isfinite(mse)
|
|
139
|
+
assert all(torch.isfinite(parameter).all() for parameter in model.parameters())
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
@pytest.mark.parametrize("model_class", [GaussianConvRBM, GaussianConvRBM4Deep])
|
|
143
|
+
def test_gaussian_conv_forward_preserves_gradient_flow(model_class):
|
|
144
|
+
model = model_class(visible_shape=(4, 4), filter_shape=(2, 2), n_filters=2)
|
|
145
|
+
with torch.no_grad():
|
|
146
|
+
model.W.fill_(0.05)
|
|
147
|
+
model.b.fill_(1)
|
|
148
|
+
samples = torch.arange(32, dtype=torch.float32).reshape(2, 1, 4, 4)
|
|
149
|
+
head = torch.nn.Linear(18, 1, bias=False)
|
|
150
|
+
with torch.no_grad():
|
|
151
|
+
head.weight.fill_(1)
|
|
152
|
+
|
|
153
|
+
head(model(samples).flatten(1)).square().mean().backward()
|
|
154
|
+
|
|
155
|
+
assert model.W.grad is not None
|
|
156
|
+
assert torch.isfinite(model.W.grad).all()
|
|
157
|
+
assert torch.count_nonzero(model.W.grad) > 0
|
|
158
|
+
|
|
159
|
+
|
|
160
|
+
def test_variance_energy_matches_marginalized_joint_distribution():
|
|
161
|
+
model = VarianceGaussianRBM(n_visible=2, n_hidden=2).double()
|
|
162
|
+
with torch.no_grad():
|
|
163
|
+
model.W.copy_(torch.tensor([[0.2, -0.1], [0.4, 0.3]]))
|
|
164
|
+
model.a.copy_(torch.tensor([0.3, -0.4]))
|
|
165
|
+
model.b.copy_(torch.tensor([-0.2, 0.1]))
|
|
166
|
+
model.sigma.copy_(torch.tensor([0.5, 2.0]))
|
|
167
|
+
samples = torch.tensor([[0.0, 1.0], [-1.0, 2.0]], dtype=torch.float64)
|
|
168
|
+
hidden = torch.tensor(
|
|
169
|
+
[[0.0, 0.0], [0.0, 1.0], [1.0, 0.0], [1.0, 1.0]], dtype=torch.float64
|
|
170
|
+
)
|
|
171
|
+
variance = model.sigma.square() + torch.finfo(samples.dtype).eps
|
|
172
|
+
quadratic = ((samples - model.a).square() / (2 * variance)).sum(1)
|
|
173
|
+
joint_energy = (
|
|
174
|
+
quadratic[:, None]
|
|
175
|
+
- hidden @ model.b
|
|
176
|
+
- (samples / variance) @ model.W @ hidden.t()
|
|
177
|
+
)
|
|
178
|
+
expected = -torch.logsumexp(-joint_energy, dim=1)
|
|
179
|
+
|
|
180
|
+
torch.testing.assert_close(model.energy(samples), expected)
|
|
181
|
+
|
|
182
|
+
|
|
183
|
+
def test_variance_visible_sampling_returns_means_and_correctly_scaled_states():
|
|
184
|
+
model = VarianceGaussianRBM(n_visible=2, n_hidden=1)
|
|
185
|
+
with torch.no_grad():
|
|
186
|
+
model.W.copy_(torch.tensor([[0.1], [0.2]]))
|
|
187
|
+
model.a.copy_(torch.tensor([0.3, -0.4]))
|
|
188
|
+
model.sigma.copy_(torch.tensor([0.5, 2.0]))
|
|
189
|
+
hidden = torch.ones(20000, 1)
|
|
190
|
+
|
|
191
|
+
torch.manual_seed(0)
|
|
192
|
+
means, states = model.visible_sampling(hidden)
|
|
193
|
+
expected_means = hidden @ model.W.t() + model.a
|
|
194
|
+
|
|
195
|
+
torch.testing.assert_close(means, expected_means)
|
|
196
|
+
torch.testing.assert_close(states.mean(0), expected_means[0], rtol=0, atol=0.04)
|
|
197
|
+
torch.testing.assert_close(states.std(0), model.sigma, rtol=0.03, atol=0)
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
def test_variance_gibbs_sampling_uses_random_visible_states():
|
|
201
|
+
model = VarianceGaussianRBM(n_visible=2, n_hidden=1)
|
|
202
|
+
with torch.no_grad():
|
|
203
|
+
model.W.zero_()
|
|
204
|
+
model.a.zero_()
|
|
205
|
+
|
|
206
|
+
torch.manual_seed(0)
|
|
207
|
+
visible_states = model.gibbs_sampling(torch.zeros(256, 2))[-1]
|
|
208
|
+
|
|
209
|
+
assert visible_states.std() > 0.5
|
|
210
|
+
|
|
211
|
+
|
|
212
|
+
def test_variance_gaussian_rbm_end_to_end():
|
|
213
|
+
torch.manual_seed(0)
|
|
214
|
+
model = VarianceGaussianRBM(n_visible=4, n_hidden=2, learning_rate=0.001)
|
|
215
|
+
dataset = TensorDataset(torch.rand(5, 4), torch.zeros(5))
|
|
216
|
+
initial_sigma = model.sigma.detach().clone()
|
|
217
|
+
|
|
218
|
+
mse, pl = model.fit(dataset, batch_size=2, epochs=2)
|
|
219
|
+
reconstruction_mse, reconstruction = model.reconstruct(dataset)
|
|
220
|
+
|
|
221
|
+
assert torch.isfinite(mse) and torch.isfinite(pl)
|
|
222
|
+
assert torch.isfinite(reconstruction_mse)
|
|
223
|
+
assert reconstruction.shape == (5, 4)
|
|
224
|
+
assert torch.isfinite(reconstruction).all()
|
|
225
|
+
assert all(torch.isfinite(parameter).all() for parameter in model.parameters())
|
|
226
|
+
assert not torch.equal(model.sigma, initial_sigma)
|
|
@@ -1,5 +1,6 @@
|
|
|
1
1
|
import logging as stdlib_logging
|
|
2
2
|
|
|
3
|
+
import matplotlib.pyplot as plt
|
|
3
4
|
import numpy as np
|
|
4
5
|
import pytest
|
|
5
6
|
import torch
|
|
@@ -73,3 +74,54 @@ def test_visual_helpers(tmp_path, monkeypatch):
|
|
|
73
74
|
|
|
74
75
|
def test_deep_sigmoid_name_remains_available():
|
|
75
76
|
assert issubclass(SigmoidRBM4Deep, SigmoidRBM)
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
@pytest.mark.parametrize("n_originals, n_reconstructed", [(1, 2), (2, 1)])
|
|
80
|
+
def test_ssim_rejects_unequal_batch_lengths(n_originals, n_reconstructed):
|
|
81
|
+
sample = torch.arange(64, dtype=torch.float32).reshape(1, 8, 8)
|
|
82
|
+
originals = sample.repeat(n_originals, 1, 1)
|
|
83
|
+
reconstructed = sample.reshape(1, 64).repeat(n_reconstructed, 1)
|
|
84
|
+
|
|
85
|
+
with pytest.raises(ValueError):
|
|
86
|
+
calculate_ssim(reconstructed, originals)
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
@pytest.fixture
|
|
90
|
+
def existing_figures():
|
|
91
|
+
figures = set(plt.get_fignums())
|
|
92
|
+
yield figures
|
|
93
|
+
for number in set(plt.get_fignums()) - figures:
|
|
94
|
+
plt.close(number)
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
@pytest.mark.parametrize("shape", [(3, 5), (1, 5, 7), (3, 5, 7)])
|
|
98
|
+
def test_tensor_render_preserves_image_layout(shape, monkeypatch, existing_figures):
|
|
99
|
+
samples = torch.rand(shape)
|
|
100
|
+
rendered = []
|
|
101
|
+
|
|
102
|
+
def capture_image():
|
|
103
|
+
rendered.append(np.asarray(plt.gca().images[0].get_array()))
|
|
104
|
+
|
|
105
|
+
monkeypatch.setattr(plt, "show", capture_image)
|
|
106
|
+
tensor.show_tensor(samples)
|
|
107
|
+
|
|
108
|
+
expected = samples.numpy()
|
|
109
|
+
if len(shape) == 3:
|
|
110
|
+
expected = np.moveaxis(expected, 0, -1) if shape[0] == 3 else expected[0]
|
|
111
|
+
np.testing.assert_array_equal(rendered[0], expected)
|
|
112
|
+
assert set(plt.get_fignums()) == existing_figures
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
@pytest.mark.parametrize("operation", ["save", "show"])
|
|
116
|
+
def test_tensor_render_closes_figure_after_failure(
|
|
117
|
+
operation, tmp_path, existing_figures
|
|
118
|
+
):
|
|
119
|
+
samples = torch.zeros(2, 5, 7)
|
|
120
|
+
|
|
121
|
+
with pytest.raises(TypeError):
|
|
122
|
+
if operation == "save":
|
|
123
|
+
tensor.save_tensor(samples, str(tmp_path / "invalid.png"))
|
|
124
|
+
else:
|
|
125
|
+
tensor.show_tensor(samples)
|
|
126
|
+
|
|
127
|
+
assert set(plt.get_fignums()) == existing_figures
|
|
@@ -1,32 +0,0 @@
|
|
|
1
|
-
"""Tensor visualization."""
|
|
2
|
-
|
|
3
|
-
import matplotlib.pyplot as plt
|
|
4
|
-
import torch
|
|
5
|
-
|
|
6
|
-
|
|
7
|
-
def _show(tensor: torch.Tensor) -> None:
|
|
8
|
-
image = tensor.permute(1, 2, 0) if tensor.size(0) == 3 else tensor
|
|
9
|
-
plt.imshow(
|
|
10
|
-
image.detach().cpu().numpy(),
|
|
11
|
-
cmap=None if tensor.size(0) == 3 else "gray",
|
|
12
|
-
)
|
|
13
|
-
plt.xticks([])
|
|
14
|
-
plt.yticks([])
|
|
15
|
-
|
|
16
|
-
|
|
17
|
-
def save_tensor(tensor: torch.Tensor, output_path: str) -> None:
|
|
18
|
-
"""Save a tensor as an image."""
|
|
19
|
-
|
|
20
|
-
plt.figure()
|
|
21
|
-
_show(tensor)
|
|
22
|
-
plt.savefig(output_path)
|
|
23
|
-
plt.close()
|
|
24
|
-
|
|
25
|
-
|
|
26
|
-
def show_tensor(tensor: torch.Tensor) -> None:
|
|
27
|
-
"""Display a tensor as an image."""
|
|
28
|
-
|
|
29
|
-
plt.figure()
|
|
30
|
-
_show(tensor)
|
|
31
|
-
plt.show()
|
|
32
|
-
plt.close()
|
|
@@ -1,73 +0,0 @@
|
|
|
1
|
-
import pytest
|
|
2
|
-
import torch
|
|
3
|
-
from torch.utils.data import TensorDataset
|
|
4
|
-
|
|
5
|
-
from learnergy.models.gaussian import (
|
|
6
|
-
GaussianConvRBM,
|
|
7
|
-
GaussianConvRBM4Deep,
|
|
8
|
-
GaussianRBM,
|
|
9
|
-
GaussianRBM4deep,
|
|
10
|
-
GaussianReluRBM,
|
|
11
|
-
GaussianReluRBM4deep,
|
|
12
|
-
GaussianSeluRBM,
|
|
13
|
-
VarianceGaussianRBM,
|
|
14
|
-
)
|
|
15
|
-
|
|
16
|
-
|
|
17
|
-
def test_gaussian_rbm_end_to_end():
|
|
18
|
-
torch.manual_seed(0)
|
|
19
|
-
dataset = TensorDataset(torch.rand(12, 16), torch.zeros(12))
|
|
20
|
-
model = GaussianRBM(n_visible=16, n_hidden=8)
|
|
21
|
-
|
|
22
|
-
mse, pl = model.fit(dataset, batch_size=4, epochs=1)
|
|
23
|
-
reconstruction_mse, reconstruction = model.reconstruct(dataset)
|
|
24
|
-
|
|
25
|
-
assert mse >= 0
|
|
26
|
-
assert torch.isfinite(pl)
|
|
27
|
-
assert reconstruction_mse >= 0
|
|
28
|
-
assert reconstruction.shape == (12, 16)
|
|
29
|
-
assert model(torch.rand(2, 16)).shape == (2, 8)
|
|
30
|
-
|
|
31
|
-
|
|
32
|
-
@pytest.mark.parametrize("model_class", [GaussianReluRBM, GaussianSeluRBM])
|
|
33
|
-
def test_gaussian_activation_variants(model_class):
|
|
34
|
-
model = model_class(n_visible=16, n_hidden=8)
|
|
35
|
-
probs, states = model.hidden_sampling(torch.rand(2, 16), scale=True)
|
|
36
|
-
|
|
37
|
-
assert probs.shape == states.shape == (2, 8)
|
|
38
|
-
assert torch.equal(probs, states)
|
|
39
|
-
|
|
40
|
-
|
|
41
|
-
def test_deep_model_names_remain_available():
|
|
42
|
-
assert issubclass(GaussianRBM4deep, GaussianRBM)
|
|
43
|
-
assert issubclass(GaussianReluRBM4deep, GaussianReluRBM)
|
|
44
|
-
assert issubclass(GaussianConvRBM4Deep, GaussianConvRBM)
|
|
45
|
-
|
|
46
|
-
|
|
47
|
-
def test_variance_gaussian_rbm_sampling_is_finite():
|
|
48
|
-
model = VarianceGaussianRBM(n_visible=16, n_hidden=8)
|
|
49
|
-
with torch.no_grad():
|
|
50
|
-
model.sigma.fill_(1e-12)
|
|
51
|
-
|
|
52
|
-
samples = torch.rand(2, 16)
|
|
53
|
-
hidden_probs, hidden_states = model.hidden_sampling(samples)
|
|
54
|
-
visible_probs, visible_states = model.visible_sampling(hidden_states)
|
|
55
|
-
|
|
56
|
-
assert torch.isfinite(hidden_probs).all()
|
|
57
|
-
assert torch.isfinite(model.energy(samples)).all()
|
|
58
|
-
assert visible_probs.shape == visible_states.shape == (2, 16)
|
|
59
|
-
assert "sigma" in model.state_dict()
|
|
60
|
-
assert len(model.optimizer.param_groups) == 2
|
|
61
|
-
|
|
62
|
-
|
|
63
|
-
def test_gaussian_conv_rbm_end_to_end():
|
|
64
|
-
dataset = TensorDataset(torch.rand(8, 1, 8, 8), torch.zeros(8))
|
|
65
|
-
model = GaussianConvRBM(
|
|
66
|
-
visible_shape=(8, 8),
|
|
67
|
-
filter_shape=(3, 3),
|
|
68
|
-
n_filters=2,
|
|
69
|
-
n_channels=1,
|
|
70
|
-
)
|
|
71
|
-
|
|
72
|
-
assert model.fit(dataset, batch_size=4, epochs=1) >= 0
|
|
73
|
-
assert model(dataset.tensors[0][:2]).shape == (2, 2, 6, 6)
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|