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.
Files changed (49) hide show
  1. {learnergy-2.0.0 → learnergy-2.0.1}/PKG-INFO +30 -4
  2. {learnergy-2.0.0 → learnergy-2.0.1}/README.md +29 -3
  3. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/__init__.py +1 -1
  4. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/math/metrics.py +13 -4
  5. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/bernoulli/discriminative_rbm.py +1 -2
  6. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/bernoulli/dropout_rbm.py +4 -28
  7. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/bernoulli/rbm.py +1 -6
  8. learnergy-2.0.1/learnergy/models/gaussian/_normalization.py +13 -0
  9. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/gaussian/gaussian_conv_rbm.py +4 -7
  10. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/gaussian/gaussian_rbm.py +18 -31
  11. learnergy-2.0.1/learnergy/visual/tensor.py +38 -0
  12. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy.egg-info/PKG-INFO +30 -4
  13. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy.egg-info/SOURCES.txt +1 -0
  14. {learnergy-2.0.0 → learnergy-2.0.1}/pyproject.toml +1 -1
  15. {learnergy-2.0.0 → learnergy-2.0.1}/tests/test_bernoulli.py +61 -0
  16. {learnergy-2.0.0 → learnergy-2.0.1}/tests/test_core.py +1 -1
  17. {learnergy-2.0.0 → learnergy-2.0.1}/tests/test_deep.py +48 -3
  18. learnergy-2.0.1/tests/test_gaussian.py +226 -0
  19. {learnergy-2.0.0 → learnergy-2.0.1}/tests/test_utilities.py +52 -0
  20. learnergy-2.0.0/learnergy/visual/tensor.py +0 -32
  21. learnergy-2.0.0/tests/test_gaussian.py +0 -73
  22. {learnergy-2.0.0 → learnergy-2.0.1}/LICENSE +0 -0
  23. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/core/__init__.py +0 -0
  24. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/core/dataset.py +0 -0
  25. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/core/model.py +0 -0
  26. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/math/__init__.py +0 -0
  27. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/math/scale.py +0 -0
  28. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/__init__.py +0 -0
  29. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/bernoulli/__init__.py +0 -0
  30. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/bernoulli/conv_rbm.py +0 -0
  31. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/bernoulli/e_dropout_rbm.py +0 -0
  32. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/deep/__init__.py +0 -0
  33. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/deep/conv_dbn.py +0 -0
  34. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/deep/dbn.py +0 -0
  35. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/deep/residual_dbn.py +0 -0
  36. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/extra/__init__.py +0 -0
  37. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/extra/sigmoid_rbm.py +0 -0
  38. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/models/gaussian/__init__.py +0 -0
  39. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/utils/__init__.py +0 -0
  40. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/utils/constants.py +0 -0
  41. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/utils/exception.py +0 -0
  42. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/utils/logging.py +0 -0
  43. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/visual/__init__.py +0 -0
  44. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/visual/convergence.py +0 -0
  45. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy/visual/image.py +0 -0
  46. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy.egg-info/dependency_links.txt +0 -0
  47. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy.egg-info/requires.txt +0 -0
  48. {learnergy-2.0.0 → learnergy-2.0.1}/learnergy.egg-info/top_level.txt +0 -0
  49. {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.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
- pip install learnergy
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
- Install the optional torchvision dependency to run the examples:
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
- pip install learnergy
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
- Install the optional torchvision dependency to run the examples:
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
@@ -2,4 +2,4 @@
2
2
  of several modules and sub-modules.
3
3
  """
4
4
 
5
- __version__ = "2.0.0"
5
+ __version__ = "2.0.1"
@@ -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
- width, height = originals.shape[1:3]
21
+ height, width = originals.shape[1:3]
13
22
 
14
23
  return sum(
15
24
  structural_similarity(
16
25
  original,
17
- rebuilt.reshape(width, height),
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
- activations = torch.exp(F.linear(h, self.U, self.c))
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
- batches = DataLoader(
127
- dataset, batch_size=batch_size, shuffle=False, num_workers=0
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).detach(), 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
- eps = torch.finfo(samples.dtype).eps
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
- eps = torch.finfo(x.dtype).eps
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
- Note that this class implements a new cost function that takes in account
425
- a new learning parameter: variance (sigma).
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
- sigma = torch.pow(self.sigma, 2) + torch.finfo(v.dtype).eps
494
- activations = F.linear(torch.div(v, sigma), self.W.t(), self.b)
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 probabilities and states of the visible layer sampling.
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
- if self.device == "cpu":
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
- sigma = torch.pow(self.sigma, 2) + torch.finfo(samples.dtype).eps
544
- activations = F.linear(torch.div(samples, sigma), self.W.t(), self.b)
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(torch.div(torch.pow(samples - self.a, 2), 2 * sigma), dim=1)
534
+ v = torch.sum((samples - self.a).square() / (2 * variance), dim=1)
548
535
 
549
- energy = -v - h
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.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
- pip install learnergy
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
- Install the optional torchvision dependency to run the examples:
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,5 +1,5 @@
1
1
  [build-system]
2
- requires = ["setuptools>=77"]
2
+ requires = ["setuptools>=84.0.0"]
3
3
  build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
@@ -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
@@ -6,7 +6,7 @@ from learnergy.core import Model
6
6
 
7
7
 
8
8
  def test_package_version():
9
- assert learnergy.__version__ == "2.0.0"
9
+ assert learnergy.__version__ == "2.0.1"
10
10
 
11
11
 
12
12
  def test_model_tracks_device_and_history():
@@ -111,8 +111,9 @@ def test_residual_dbn_validates_weights():
111
111
  ResidualDBN(zetta2=-1)
112
112
 
113
113
 
114
- def test_conv_dbn_builds_trains_and_reconstructs():
115
- dataset = TensorDataset(torch.rand(8, 1, 8, 8), torch.zeros(8))
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 == (8, 1, 8, 8)
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