learnergy 1.2.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 (74) hide show
  1. learnergy-2.0.1/PKG-INFO +168 -0
  2. learnergy-2.0.1/README.md +131 -0
  3. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/__init__.py +1 -1
  4. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/core/__init__.py +3 -2
  5. learnergy-2.0.1/learnergy/core/dataset.py +48 -0
  6. learnergy-2.0.1/learnergy/core/model.py +58 -0
  7. learnergy-2.0.1/learnergy/math/__init__.py +1 -0
  8. learnergy-2.0.1/learnergy/math/metrics.py +30 -0
  9. learnergy-2.0.1/learnergy/math/scale.py +14 -0
  10. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/__init__.py +1 -2
  11. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/bernoulli/__init__.py +12 -3
  12. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/bernoulli/conv_rbm.py +71 -224
  13. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/bernoulli/discriminative_rbm.py +39 -127
  14. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/bernoulli/dropout_rbm.py +16 -104
  15. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/bernoulli/e_dropout_rbm.py +24 -50
  16. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/bernoulli/rbm.py +47 -191
  17. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/deep/__init__.py +3 -2
  18. learnergy-2.0.1/learnergy/models/deep/conv_dbn.py +247 -0
  19. learnergy-2.0.1/learnergy/models/deep/dbn.py +231 -0
  20. learnergy-2.0.1/learnergy/models/deep/residual_dbn.py +130 -0
  21. learnergy-2.0.1/learnergy/models/extra/__init__.py +5 -0
  22. learnergy-2.0.1/learnergy/models/extra/sigmoid_rbm.py +34 -0
  23. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/gaussian/__init__.py +12 -2
  24. learnergy-2.0.1/learnergy/models/gaussian/_normalization.py +13 -0
  25. learnergy-2.0.1/learnergy/models/gaussian/gaussian_conv_rbm.py +128 -0
  26. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/gaussian/gaussian_rbm.py +45 -331
  27. learnergy-2.0.1/learnergy/utils/__init__.py +1 -0
  28. learnergy-2.0.1/learnergy/utils/constants.py +3 -0
  29. learnergy-2.0.1/learnergy/utils/exception.py +40 -0
  30. learnergy-2.0.1/learnergy/utils/logging.py +37 -0
  31. learnergy-2.0.1/learnergy/visual/__init__.py +1 -0
  32. learnergy-2.0.1/learnergy/visual/convergence.py +41 -0
  33. learnergy-2.0.1/learnergy/visual/image.py +97 -0
  34. learnergy-2.0.1/learnergy/visual/tensor.py +38 -0
  35. learnergy-2.0.1/learnergy.egg-info/PKG-INFO +168 -0
  36. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy.egg-info/SOURCES.txt +6 -3
  37. learnergy-2.0.1/learnergy.egg-info/requires.txt +18 -0
  38. learnergy-2.0.1/pyproject.toml +71 -0
  39. learnergy-2.0.1/setup.cfg +4 -0
  40. learnergy-2.0.1/tests/test_bernoulli.py +234 -0
  41. learnergy-2.0.1/tests/test_core.py +23 -0
  42. learnergy-2.0.1/tests/test_deep.py +207 -0
  43. learnergy-2.0.1/tests/test_gaussian.py +226 -0
  44. learnergy-2.0.1/tests/test_utilities.py +127 -0
  45. learnergy-1.2.0/PKG-INFO +0 -193
  46. learnergy-1.2.0/README.md +0 -146
  47. learnergy-1.2.0/learnergy/core/dataset.py +0 -109
  48. learnergy-1.2.0/learnergy/core/model.py +0 -72
  49. learnergy-1.2.0/learnergy/math/__init__.py +0 -2
  50. learnergy-1.2.0/learnergy/math/metrics.py +0 -40
  51. learnergy-1.2.0/learnergy/math/scale.py +0 -25
  52. learnergy-1.2.0/learnergy/models/deep/conv_dbn.py +0 -435
  53. learnergy-1.2.0/learnergy/models/deep/dbn.py +0 -407
  54. learnergy-1.2.0/learnergy/models/deep/residual_dbn.py +0 -205
  55. learnergy-1.2.0/learnergy/models/extra/__init__.py +0 -4
  56. learnergy-1.2.0/learnergy/models/extra/sigmoid_rbm.py +0 -207
  57. learnergy-1.2.0/learnergy/models/gaussian/gaussian_conv_rbm.py +0 -430
  58. learnergy-1.2.0/learnergy/utils/__init__.py +0 -2
  59. learnergy-1.2.0/learnergy/utils/constants.py +0 -6
  60. learnergy-1.2.0/learnergy/utils/exception.py +0 -99
  61. learnergy-1.2.0/learnergy/utils/logging.py +0 -60
  62. learnergy-1.2.0/learnergy/visual/__init__.py +0 -2
  63. learnergy-1.2.0/learnergy/visual/convergence.py +0 -63
  64. learnergy-1.2.0/learnergy/visual/image.py +0 -141
  65. learnergy-1.2.0/learnergy/visual/tensor.py +0 -52
  66. learnergy-1.2.0/learnergy.egg-info/PKG-INFO +0 -193
  67. learnergy-1.2.0/learnergy.egg-info/requires.txt +0 -16
  68. learnergy-1.2.0/pyproject.toml +0 -2
  69. learnergy-1.2.0/setup.cfg +0 -14
  70. learnergy-1.2.0/setup.py +0 -50
  71. learnergy-1.2.0/tests/test_bugfixes.py +0 -363
  72. {learnergy-1.2.0 → learnergy-2.0.1}/LICENSE +0 -0
  73. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy.egg-info/dependency_links.txt +0 -0
  74. {learnergy-1.2.0 → learnergy-2.0.1}/learnergy.egg-info/top_level.txt +0 -0
@@ -0,0 +1,168 @@
1
+ Metadata-Version: 2.4
2
+ Name: learnergy
3
+ Version: 2.0.1
4
+ Summary: Energy-based machine learners built with PyTorch
5
+ Author-email: Mateus Roder <mateus.roder@unesp.br>, Gustavo de Rosa <gustavo.rosa@unesp.br>
6
+ License-Expression: Apache-2.0
7
+ Project-URL: Homepage, https://github.com/gugarosa/learnergy
8
+ Project-URL: Documentation, https://learnergy.readthedocs.io
9
+ Project-URL: Issues, https://github.com/gugarosa/learnergy/issues
10
+ Classifier: Development Status :: 5 - Production/Stable
11
+ Classifier: Intended Audience :: Developers
12
+ Classifier: Intended Audience :: Education
13
+ Classifier: Intended Audience :: Science/Research
14
+ Classifier: Programming Language :: Python :: 3.11
15
+ Classifier: Programming Language :: Python :: 3.12
16
+ Classifier: Programming Language :: Python :: 3.13
17
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
18
+ Classifier: Topic :: Software Development :: Libraries :: Python Modules
19
+ Requires-Python: >=3.11
20
+ Description-Content-Type: text/markdown
21
+ License-File: LICENSE
22
+ Requires-Dist: matplotlib>=3.10.9
23
+ Requires-Dist: numpy>=2.0
24
+ Requires-Dist: Pillow>=8.1.2
25
+ Requires-Dist: scikit-image>=0.26.0
26
+ Requires-Dist: torch>=2.13.0
27
+ Provides-Extra: dev
28
+ Requires-Dist: pre-commit>=4.6.0; extra == "dev"
29
+ Provides-Extra: docs
30
+ Requires-Dist: sphinx>=9; extra == "docs"
31
+ Provides-Extra: examples
32
+ Requires-Dist: torchvision>=0.9.0; extra == "examples"
33
+ Provides-Extra: tests
34
+ Requires-Dist: coverage>=7.10; extra == "tests"
35
+ Requires-Dist: pytest>=9.0.2; extra == "tests"
36
+ Dynamic: license-file
37
+
38
+ # Learnergy: Energy-based Machine Learners
39
+
40
+ [![Latest release](https://img.shields.io/github/release/gugarosa/learnergy.svg)](https://github.com/gugarosa/learnergy/releases)
41
+ [![CI](https://github.com/gugarosa/learnergy/actions/workflows/ci.yml/badge.svg)](https://github.com/gugarosa/learnergy/actions/workflows/ci.yml)
42
+ [![DOI](https://img.shields.io/badge/DOI-10.5281/zenodo.4390744-006DB9.svg)](https://doi.org/10.5281/zenodo.4390744)
43
+ [![License](https://img.shields.io/github/license/gugarosa/learnergy.svg)](LICENSE)
44
+
45
+ Learnergy provides PyTorch implementations of Restricted Boltzmann Machines
46
+ (RBMs) and Deep Belief Networks (DBNs) for unsupervised feature learning,
47
+ generative modeling, and classification. It also includes dataset adapters,
48
+ image-quality metrics, and visualization helpers.
49
+
50
+ ## Installation
51
+
52
+ Learnergy requires Python 3.11 or newer. Add it to a project managed by uv with:
53
+
54
+ ```bash
55
+ uv add learnergy
56
+ ```
57
+
58
+ Add the optional torchvision dependency to run the examples:
59
+
60
+ ```bash
61
+ uv add "learnergy[examples]"
62
+ ```
63
+
64
+ For a consumer installation in an existing Python environment, pip is also supported:
65
+
66
+ ```bash
67
+ pip install learnergy
68
+ pip install "learnergy[examples]"
69
+ ```
70
+
71
+ ## Quick start
72
+
73
+ ```python
74
+ import torch
75
+ from torch.utils.data import TensorDataset
76
+
77
+ from learnergy.models.bernoulli import RBM
78
+
79
+ samples = torch.bernoulli(torch.rand(1_024, 784))
80
+ targets = torch.zeros(1_024)
81
+ dataset = TensorDataset(samples, targets)
82
+
83
+ model = RBM(n_visible=784, n_hidden=128, learning_rate=0.1)
84
+ mse, pseudo_likelihood = model.fit(dataset, batch_size=128, epochs=5)
85
+ reconstruction_mse, reconstructed = model.reconstruct(dataset)
86
+ ```
87
+
88
+ Stack RBMs into a DBN:
89
+
90
+ ```python
91
+ from learnergy.models.deep import DBN
92
+
93
+ model = DBN(
94
+ model=("gaussian", "sigmoid"),
95
+ n_visible=784,
96
+ n_hidden=(256, 128),
97
+ steps=(1, 1),
98
+ learning_rate=(0.01, 0.01),
99
+ momentum=(0, 0),
100
+ decay=(0, 0),
101
+ temperature=(1, 1),
102
+ )
103
+ model.fit(dataset, batch_size=128, epochs=(5, 5))
104
+ ```
105
+
106
+ ## Available models
107
+
108
+ | Family | Models |
109
+ |---|---|
110
+ | Bernoulli | `RBM`, `ConvRBM`, `DiscriminativeRBM`, `HybridDiscriminativeRBM`, `DropoutRBM`, `DropConnectRBM`, `EDropoutRBM` |
111
+ | Gaussian | `GaussianRBM`, `GaussianReluRBM`, `GaussianSeluRBM`, `VarianceGaussianRBM`, `GaussianConvRBM` |
112
+ | Extra | `SigmoidRBM` |
113
+ | Deep | `DBN`, `ConvDBN`, `ResidualDBN` |
114
+
115
+ The `learnergy.core.Dataset`, `learnergy.math`, and `learnergy.visual` modules
116
+ remain available for array-backed datasets, SSIM/scaling helpers, convergence
117
+ plots, image mosaics, and tensor rendering.
118
+
119
+ See [`examples/applications`](examples/applications) for complete training and
120
+ classification programs.
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
+
141
+ ## Development
142
+
143
+ The repository uses [uv](https://docs.astral.sh/uv/) for reproducible
144
+ environments and packaging:
145
+
146
+ ```bash
147
+ uv sync --locked
148
+ uv run pytest
149
+ uv build
150
+ ```
151
+
152
+ ## Citation
153
+
154
+ ```bibtex
155
+ @misc{roder2020learnergy,
156
+ title={Learnergy: Energy-based Machine Learners},
157
+ author={Mateus Roder and Gustavo Henrique de Rosa and João Paulo Papa},
158
+ year={2020},
159
+ eprint={2003.07443},
160
+ archivePrefix={arXiv},
161
+ primaryClass={cs.LG}
162
+ }
163
+ ```
164
+
165
+ ## Support
166
+
167
+ Open an [issue](https://github.com/gugarosa/learnergy/issues) for bug reports
168
+ and questions.
@@ -0,0 +1,131 @@
1
+ # Learnergy: Energy-based Machine Learners
2
+
3
+ [![Latest release](https://img.shields.io/github/release/gugarosa/learnergy.svg)](https://github.com/gugarosa/learnergy/releases)
4
+ [![CI](https://github.com/gugarosa/learnergy/actions/workflows/ci.yml/badge.svg)](https://github.com/gugarosa/learnergy/actions/workflows/ci.yml)
5
+ [![DOI](https://img.shields.io/badge/DOI-10.5281/zenodo.4390744-006DB9.svg)](https://doi.org/10.5281/zenodo.4390744)
6
+ [![License](https://img.shields.io/github/license/gugarosa/learnergy.svg)](LICENSE)
7
+
8
+ Learnergy provides PyTorch implementations of Restricted Boltzmann Machines
9
+ (RBMs) and Deep Belief Networks (DBNs) for unsupervised feature learning,
10
+ generative modeling, and classification. It also includes dataset adapters,
11
+ image-quality metrics, and visualization helpers.
12
+
13
+ ## Installation
14
+
15
+ Learnergy requires Python 3.11 or newer. Add it to a project managed by uv with:
16
+
17
+ ```bash
18
+ uv add learnergy
19
+ ```
20
+
21
+ Add the optional torchvision dependency to run the examples:
22
+
23
+ ```bash
24
+ uv add "learnergy[examples]"
25
+ ```
26
+
27
+ For a consumer installation in an existing Python environment, pip is also supported:
28
+
29
+ ```bash
30
+ pip install learnergy
31
+ pip install "learnergy[examples]"
32
+ ```
33
+
34
+ ## Quick start
35
+
36
+ ```python
37
+ import torch
38
+ from torch.utils.data import TensorDataset
39
+
40
+ from learnergy.models.bernoulli import RBM
41
+
42
+ samples = torch.bernoulli(torch.rand(1_024, 784))
43
+ targets = torch.zeros(1_024)
44
+ dataset = TensorDataset(samples, targets)
45
+
46
+ model = RBM(n_visible=784, n_hidden=128, learning_rate=0.1)
47
+ mse, pseudo_likelihood = model.fit(dataset, batch_size=128, epochs=5)
48
+ reconstruction_mse, reconstructed = model.reconstruct(dataset)
49
+ ```
50
+
51
+ Stack RBMs into a DBN:
52
+
53
+ ```python
54
+ from learnergy.models.deep import DBN
55
+
56
+ model = DBN(
57
+ model=("gaussian", "sigmoid"),
58
+ n_visible=784,
59
+ n_hidden=(256, 128),
60
+ steps=(1, 1),
61
+ learning_rate=(0.01, 0.01),
62
+ momentum=(0, 0),
63
+ decay=(0, 0),
64
+ temperature=(1, 1),
65
+ )
66
+ model.fit(dataset, batch_size=128, epochs=(5, 5))
67
+ ```
68
+
69
+ ## Available models
70
+
71
+ | Family | Models |
72
+ |---|---|
73
+ | Bernoulli | `RBM`, `ConvRBM`, `DiscriminativeRBM`, `HybridDiscriminativeRBM`, `DropoutRBM`, `DropConnectRBM`, `EDropoutRBM` |
74
+ | Gaussian | `GaussianRBM`, `GaussianReluRBM`, `GaussianSeluRBM`, `VarianceGaussianRBM`, `GaussianConvRBM` |
75
+ | Extra | `SigmoidRBM` |
76
+ | Deep | `DBN`, `ConvDBN`, `ResidualDBN` |
77
+
78
+ The `learnergy.core.Dataset`, `learnergy.math`, and `learnergy.visual` modules
79
+ remain available for array-backed datasets, SSIM/scaling helpers, convergence
80
+ plots, image mosaics, and tensor rendering.
81
+
82
+ See [`examples/applications`](examples/applications) for complete training and
83
+ classification programs.
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
+
104
+ ## Development
105
+
106
+ The repository uses [uv](https://docs.astral.sh/uv/) for reproducible
107
+ environments and packaging:
108
+
109
+ ```bash
110
+ uv sync --locked
111
+ uv run pytest
112
+ uv build
113
+ ```
114
+
115
+ ## Citation
116
+
117
+ ```bibtex
118
+ @misc{roder2020learnergy,
119
+ title={Learnergy: Energy-based Machine Learners},
120
+ author={Mateus Roder and Gustavo Henrique de Rosa and João Paulo Papa},
121
+ year={2020},
122
+ eprint={2003.07443},
123
+ archivePrefix={arXiv},
124
+ primaryClass={cs.LG}
125
+ }
126
+ ```
127
+
128
+ ## Support
129
+
130
+ Open an [issue](https://github.com/gugarosa/learnergy/issues) for bug reports
131
+ and questions.
@@ -2,4 +2,4 @@
2
2
  of several modules and sub-modules.
3
3
  """
4
4
 
5
- __version__ = "1.2.0"
5
+ __version__ = "2.0.1"
@@ -1,5 +1,6 @@
1
- """A core package for all common learnergy modules.
2
- """
1
+ """Core model primitives."""
3
2
 
4
3
  from learnergy.core.dataset import Dataset
5
4
  from learnergy.core.model import Model
5
+
6
+ __all__ = ["Dataset", "Model"]
@@ -0,0 +1,48 @@
1
+ """Dataset helpers."""
2
+
3
+ from collections.abc import Callable
4
+
5
+ import torch
6
+
7
+ import learnergy.utils.exception as e
8
+ from learnergy.core.model import _validated_property
9
+ from learnergy.utils import logging
10
+
11
+ logger = logging.get_logger(__name__)
12
+
13
+
14
+ class Dataset(torch.utils.data.Dataset):
15
+ """Wrap samples and targets with an optional sample transform."""
16
+
17
+ data = _validated_property("data")
18
+ targets = _validated_property("targets")
19
+ transform = _validated_property(
20
+ "transform",
21
+ lambda _, value: value is None or callable(value),
22
+ e.TypeError,
23
+ "`transform` should be a callable or None",
24
+ )
25
+
26
+ def __init__(
27
+ self,
28
+ data,
29
+ targets,
30
+ transform: Callable | None = None,
31
+ show_log: bool = True,
32
+ ) -> None:
33
+ self.data = data
34
+ self.targets = targets
35
+ self.transform = transform
36
+
37
+ if show_log:
38
+ logger.info("Creating class: Dataset.")
39
+ logger.info("Class created.")
40
+
41
+ def __getitem__(self, idx: int):
42
+ sample = self.data[idx]
43
+ if self.transform:
44
+ sample = self.transform(sample)
45
+ return sample, self.targets[idx]
46
+
47
+ def __len__(self) -> int:
48
+ return len(self.data)
@@ -0,0 +1,58 @@
1
+ """Standard model-related implementation."""
2
+
3
+ from collections.abc import Callable
4
+ from typing import Any
5
+
6
+ import torch
7
+
8
+ import learnergy.utils.exception as e
9
+
10
+
11
+ def _validated_property(
12
+ name: str,
13
+ validator: Callable[[Any, Any], bool] | None = None,
14
+ error: type[Exception] = ValueError,
15
+ message: str = "invalid value",
16
+ ) -> property:
17
+ storage_name = f"_{name}"
18
+
19
+ def getter(instance):
20
+ return getattr(instance, storage_name)
21
+
22
+ def setter(instance, value):
23
+ if validator is not None and not validator(instance, value):
24
+ raise error(message)
25
+ setattr(instance, storage_name, value)
26
+
27
+ return property(getter, setter)
28
+
29
+
30
+ class Model(torch.nn.Module):
31
+ """Base class for Learnergy models."""
32
+
33
+ device = _validated_property(
34
+ "device",
35
+ lambda _, value: value in ("cpu", "cuda"),
36
+ e.TypeError,
37
+ "`device` should be `cpu` or `cuda`",
38
+ )
39
+ history = _validated_property("history")
40
+
41
+ def __init__(self, use_gpu: bool = False) -> None:
42
+ """Initialization method.
43
+
44
+ Args:
45
+ use_gpu: Whether GPU should be used or not.
46
+
47
+ """
48
+
49
+ super().__init__()
50
+ torch.set_default_dtype(torch.float32)
51
+ self.device = "cuda" if use_gpu and torch.cuda.is_available() else "cpu"
52
+ self.history = {}
53
+
54
+ def dump(self, **kwargs) -> None:
55
+ """Dumps any amount of keyword documents to lists in the history property."""
56
+
57
+ for k, v in kwargs.items():
58
+ self.history.setdefault(k, []).append(v)
@@ -0,0 +1 @@
1
+ """Mathematical helpers."""
@@ -0,0 +1,30 @@
1
+ """Image similarity metrics."""
2
+
3
+ import torch
4
+ from skimage.metrics import structural_similarity
5
+
6
+
7
+ def calculate_ssim(v: torch.Tensor, x: torch.Tensor) -> float:
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
+ """
18
+
19
+ originals = x.detach().cpu().numpy()
20
+ reconstructed = v.detach().cpu().numpy()
21
+ height, width = originals.shape[1:3]
22
+
23
+ return sum(
24
+ structural_similarity(
25
+ original,
26
+ rebuilt.reshape(height, width),
27
+ data_range=original.max() - original.min(),
28
+ )
29
+ for original, rebuilt in zip(originals, reconstructed, strict=True)
30
+ ) / len(reconstructed)
@@ -0,0 +1,14 @@
1
+ """Scaling helpers."""
2
+
3
+ import numpy as np
4
+
5
+ from learnergy.utils.constants import EPSILON
6
+
7
+
8
+ def unitary_scale(x: np.ndarray) -> np.ndarray:
9
+ """Scale an array to the interval from zero to one."""
10
+
11
+ scaled = x.astype("float32")
12
+ scaled -= scaled.min()
13
+ scaled /= scaled.max() + EPSILON
14
+ return scaled
@@ -1,2 +1 @@
1
- """A package contaning subpackages of models for all common learnergy modules.
2
- """
1
+ """A package contaning subpackages of models for all common learnergy modules."""
@@ -1,7 +1,5 @@
1
- """A package contaning bernoulli-based models (networks) for all common learnergy modules.
2
- """
1
+ """Bernoulli-valued RBM variants."""
3
2
 
4
- from learnergy.models.bernoulli.rbm import RBM
5
3
  from learnergy.models.bernoulli.conv_rbm import ConvRBM
6
4
  from learnergy.models.bernoulli.discriminative_rbm import (
7
5
  DiscriminativeRBM,
@@ -9,3 +7,14 @@ from learnergy.models.bernoulli.discriminative_rbm import (
9
7
  )
10
8
  from learnergy.models.bernoulli.dropout_rbm import DropConnectRBM, DropoutRBM
11
9
  from learnergy.models.bernoulli.e_dropout_rbm import EDropoutRBM
10
+ from learnergy.models.bernoulli.rbm import RBM
11
+
12
+ __all__ = [
13
+ "ConvRBM",
14
+ "DiscriminativeRBM",
15
+ "DropConnectRBM",
16
+ "DropoutRBM",
17
+ "EDropoutRBM",
18
+ "HybridDiscriminativeRBM",
19
+ "RBM",
20
+ ]