learnergy 1.2.0__tar.gz → 2.0.0__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/PKG-INFO +142 -0
- learnergy-2.0.0/README.md +105 -0
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy/__init__.py +1 -1
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy/core/__init__.py +3 -2
- learnergy-2.0.0/learnergy/core/dataset.py +48 -0
- learnergy-2.0.0/learnergy/core/model.py +58 -0
- learnergy-2.0.0/learnergy/math/__init__.py +1 -0
- learnergy-2.0.0/learnergy/math/metrics.py +21 -0
- learnergy-2.0.0/learnergy/math/scale.py +14 -0
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy/models/__init__.py +1 -2
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy/models/bernoulli/__init__.py +12 -3
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy/models/bernoulli/conv_rbm.py +71 -224
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy/models/bernoulli/discriminative_rbm.py +38 -125
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy/models/bernoulli/dropout_rbm.py +14 -78
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy/models/bernoulli/e_dropout_rbm.py +24 -50
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy/models/bernoulli/rbm.py +50 -189
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy/models/deep/__init__.py +3 -2
- learnergy-2.0.0/learnergy/models/deep/conv_dbn.py +247 -0
- learnergy-2.0.0/learnergy/models/deep/dbn.py +231 -0
- learnergy-2.0.0/learnergy/models/deep/residual_dbn.py +130 -0
- learnergy-2.0.0/learnergy/models/extra/__init__.py +5 -0
- learnergy-2.0.0/learnergy/models/extra/sigmoid_rbm.py +34 -0
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy/models/gaussian/__init__.py +12 -2
- learnergy-2.0.0/learnergy/models/gaussian/gaussian_conv_rbm.py +131 -0
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy/models/gaussian/gaussian_rbm.py +36 -309
- learnergy-2.0.0/learnergy/utils/__init__.py +1 -0
- learnergy-2.0.0/learnergy/utils/constants.py +3 -0
- learnergy-2.0.0/learnergy/utils/exception.py +40 -0
- learnergy-2.0.0/learnergy/utils/logging.py +37 -0
- learnergy-2.0.0/learnergy/visual/__init__.py +1 -0
- learnergy-2.0.0/learnergy/visual/convergence.py +41 -0
- learnergy-2.0.0/learnergy/visual/image.py +97 -0
- learnergy-2.0.0/learnergy/visual/tensor.py +32 -0
- learnergy-2.0.0/learnergy.egg-info/PKG-INFO +142 -0
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy.egg-info/SOURCES.txt +5 -3
- learnergy-2.0.0/learnergy.egg-info/requires.txt +18 -0
- learnergy-2.0.0/pyproject.toml +71 -0
- learnergy-2.0.0/setup.cfg +4 -0
- learnergy-2.0.0/tests/test_bernoulli.py +173 -0
- learnergy-2.0.0/tests/test_core.py +23 -0
- learnergy-2.0.0/tests/test_deep.py +162 -0
- learnergy-2.0.0/tests/test_gaussian.py +73 -0
- learnergy-2.0.0/tests/test_utilities.py +75 -0
- learnergy-1.2.0/PKG-INFO +0 -193
- learnergy-1.2.0/README.md +0 -146
- learnergy-1.2.0/learnergy/core/dataset.py +0 -109
- learnergy-1.2.0/learnergy/core/model.py +0 -72
- learnergy-1.2.0/learnergy/math/__init__.py +0 -2
- learnergy-1.2.0/learnergy/math/metrics.py +0 -40
- learnergy-1.2.0/learnergy/math/scale.py +0 -25
- learnergy-1.2.0/learnergy/models/deep/conv_dbn.py +0 -435
- learnergy-1.2.0/learnergy/models/deep/dbn.py +0 -407
- learnergy-1.2.0/learnergy/models/deep/residual_dbn.py +0 -205
- learnergy-1.2.0/learnergy/models/extra/__init__.py +0 -4
- learnergy-1.2.0/learnergy/models/extra/sigmoid_rbm.py +0 -207
- learnergy-1.2.0/learnergy/models/gaussian/gaussian_conv_rbm.py +0 -430
- learnergy-1.2.0/learnergy/utils/__init__.py +0 -2
- learnergy-1.2.0/learnergy/utils/constants.py +0 -6
- learnergy-1.2.0/learnergy/utils/exception.py +0 -99
- learnergy-1.2.0/learnergy/utils/logging.py +0 -60
- learnergy-1.2.0/learnergy/visual/__init__.py +0 -2
- learnergy-1.2.0/learnergy/visual/convergence.py +0 -63
- learnergy-1.2.0/learnergy/visual/image.py +0 -141
- learnergy-1.2.0/learnergy/visual/tensor.py +0 -52
- learnergy-1.2.0/learnergy.egg-info/PKG-INFO +0 -193
- learnergy-1.2.0/learnergy.egg-info/requires.txt +0 -16
- learnergy-1.2.0/pyproject.toml +0 -2
- learnergy-1.2.0/setup.cfg +0 -14
- learnergy-1.2.0/setup.py +0 -50
- learnergy-1.2.0/tests/test_bugfixes.py +0 -363
- {learnergy-1.2.0 → learnergy-2.0.0}/LICENSE +0 -0
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy.egg-info/dependency_links.txt +0 -0
- {learnergy-1.2.0 → learnergy-2.0.0}/learnergy.egg-info/top_level.txt +0 -0
learnergy-2.0.0/PKG-INFO
ADDED
|
@@ -0,0 +1,142 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: learnergy
|
|
3
|
+
Version: 2.0.0
|
|
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
|
+
[](https://github.com/gugarosa/learnergy/releases)
|
|
41
|
+
[](https://github.com/gugarosa/learnergy/actions/workflows/ci.yml)
|
|
42
|
+
[](https://doi.org/10.5281/zenodo.4390744)
|
|
43
|
+
[](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.
|
|
53
|
+
|
|
54
|
+
```bash
|
|
55
|
+
pip install learnergy
|
|
56
|
+
```
|
|
57
|
+
|
|
58
|
+
Install the optional torchvision dependency to run the examples:
|
|
59
|
+
|
|
60
|
+
```bash
|
|
61
|
+
pip install "learnergy[examples]"
|
|
62
|
+
```
|
|
63
|
+
|
|
64
|
+
## Quick start
|
|
65
|
+
|
|
66
|
+
```python
|
|
67
|
+
import torch
|
|
68
|
+
from torch.utils.data import TensorDataset
|
|
69
|
+
|
|
70
|
+
from learnergy.models.bernoulli import RBM
|
|
71
|
+
|
|
72
|
+
samples = torch.bernoulli(torch.rand(1_024, 784))
|
|
73
|
+
targets = torch.zeros(1_024)
|
|
74
|
+
dataset = TensorDataset(samples, targets)
|
|
75
|
+
|
|
76
|
+
model = RBM(n_visible=784, n_hidden=128, learning_rate=0.1)
|
|
77
|
+
mse, pseudo_likelihood = model.fit(dataset, batch_size=128, epochs=5)
|
|
78
|
+
reconstruction_mse, reconstructed = model.reconstruct(dataset)
|
|
79
|
+
```
|
|
80
|
+
|
|
81
|
+
Stack RBMs into a DBN:
|
|
82
|
+
|
|
83
|
+
```python
|
|
84
|
+
from learnergy.models.deep import DBN
|
|
85
|
+
|
|
86
|
+
model = DBN(
|
|
87
|
+
model=("gaussian", "sigmoid"),
|
|
88
|
+
n_visible=784,
|
|
89
|
+
n_hidden=(256, 128),
|
|
90
|
+
steps=(1, 1),
|
|
91
|
+
learning_rate=(0.01, 0.01),
|
|
92
|
+
momentum=(0, 0),
|
|
93
|
+
decay=(0, 0),
|
|
94
|
+
temperature=(1, 1),
|
|
95
|
+
)
|
|
96
|
+
model.fit(dataset, batch_size=128, epochs=(5, 5))
|
|
97
|
+
```
|
|
98
|
+
|
|
99
|
+
## Available models
|
|
100
|
+
|
|
101
|
+
| Family | Models |
|
|
102
|
+
|---|---|
|
|
103
|
+
| Bernoulli | `RBM`, `ConvRBM`, `DiscriminativeRBM`, `HybridDiscriminativeRBM`, `DropoutRBM`, `DropConnectRBM`, `EDropoutRBM` |
|
|
104
|
+
| Gaussian | `GaussianRBM`, `GaussianReluRBM`, `GaussianSeluRBM`, `VarianceGaussianRBM`, `GaussianConvRBM` |
|
|
105
|
+
| Extra | `SigmoidRBM` |
|
|
106
|
+
| Deep | `DBN`, `ConvDBN`, `ResidualDBN` |
|
|
107
|
+
|
|
108
|
+
The `learnergy.core.Dataset`, `learnergy.math`, and `learnergy.visual` modules
|
|
109
|
+
remain available for array-backed datasets, SSIM/scaling helpers, convergence
|
|
110
|
+
plots, image mosaics, and tensor rendering.
|
|
111
|
+
|
|
112
|
+
See [`examples/applications`](examples/applications) for complete training and
|
|
113
|
+
classification programs.
|
|
114
|
+
|
|
115
|
+
## Development
|
|
116
|
+
|
|
117
|
+
The repository uses [uv](https://docs.astral.sh/uv/) for reproducible
|
|
118
|
+
environments and packaging:
|
|
119
|
+
|
|
120
|
+
```bash
|
|
121
|
+
uv sync --locked
|
|
122
|
+
uv run pytest
|
|
123
|
+
uv build
|
|
124
|
+
```
|
|
125
|
+
|
|
126
|
+
## Citation
|
|
127
|
+
|
|
128
|
+
```bibtex
|
|
129
|
+
@misc{roder2020learnergy,
|
|
130
|
+
title={Learnergy: Energy-based Machine Learners},
|
|
131
|
+
author={Mateus Roder and Gustavo Henrique de Rosa and João Paulo Papa},
|
|
132
|
+
year={2020},
|
|
133
|
+
eprint={2003.07443},
|
|
134
|
+
archivePrefix={arXiv},
|
|
135
|
+
primaryClass={cs.LG}
|
|
136
|
+
}
|
|
137
|
+
```
|
|
138
|
+
|
|
139
|
+
## Support
|
|
140
|
+
|
|
141
|
+
Open an [issue](https://github.com/gugarosa/learnergy/issues) for bug reports
|
|
142
|
+
and questions.
|
|
@@ -0,0 +1,105 @@
|
|
|
1
|
+
# Learnergy: Energy-based Machine Learners
|
|
2
|
+
|
|
3
|
+
[](https://github.com/gugarosa/learnergy/releases)
|
|
4
|
+
[](https://github.com/gugarosa/learnergy/actions/workflows/ci.yml)
|
|
5
|
+
[](https://doi.org/10.5281/zenodo.4390744)
|
|
6
|
+
[](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.
|
|
16
|
+
|
|
17
|
+
```bash
|
|
18
|
+
pip install learnergy
|
|
19
|
+
```
|
|
20
|
+
|
|
21
|
+
Install the optional torchvision dependency to run the examples:
|
|
22
|
+
|
|
23
|
+
```bash
|
|
24
|
+
pip install "learnergy[examples]"
|
|
25
|
+
```
|
|
26
|
+
|
|
27
|
+
## Quick start
|
|
28
|
+
|
|
29
|
+
```python
|
|
30
|
+
import torch
|
|
31
|
+
from torch.utils.data import TensorDataset
|
|
32
|
+
|
|
33
|
+
from learnergy.models.bernoulli import RBM
|
|
34
|
+
|
|
35
|
+
samples = torch.bernoulli(torch.rand(1_024, 784))
|
|
36
|
+
targets = torch.zeros(1_024)
|
|
37
|
+
dataset = TensorDataset(samples, targets)
|
|
38
|
+
|
|
39
|
+
model = RBM(n_visible=784, n_hidden=128, learning_rate=0.1)
|
|
40
|
+
mse, pseudo_likelihood = model.fit(dataset, batch_size=128, epochs=5)
|
|
41
|
+
reconstruction_mse, reconstructed = model.reconstruct(dataset)
|
|
42
|
+
```
|
|
43
|
+
|
|
44
|
+
Stack RBMs into a DBN:
|
|
45
|
+
|
|
46
|
+
```python
|
|
47
|
+
from learnergy.models.deep import DBN
|
|
48
|
+
|
|
49
|
+
model = DBN(
|
|
50
|
+
model=("gaussian", "sigmoid"),
|
|
51
|
+
n_visible=784,
|
|
52
|
+
n_hidden=(256, 128),
|
|
53
|
+
steps=(1, 1),
|
|
54
|
+
learning_rate=(0.01, 0.01),
|
|
55
|
+
momentum=(0, 0),
|
|
56
|
+
decay=(0, 0),
|
|
57
|
+
temperature=(1, 1),
|
|
58
|
+
)
|
|
59
|
+
model.fit(dataset, batch_size=128, epochs=(5, 5))
|
|
60
|
+
```
|
|
61
|
+
|
|
62
|
+
## Available models
|
|
63
|
+
|
|
64
|
+
| Family | Models |
|
|
65
|
+
|---|---|
|
|
66
|
+
| Bernoulli | `RBM`, `ConvRBM`, `DiscriminativeRBM`, `HybridDiscriminativeRBM`, `DropoutRBM`, `DropConnectRBM`, `EDropoutRBM` |
|
|
67
|
+
| Gaussian | `GaussianRBM`, `GaussianReluRBM`, `GaussianSeluRBM`, `VarianceGaussianRBM`, `GaussianConvRBM` |
|
|
68
|
+
| Extra | `SigmoidRBM` |
|
|
69
|
+
| Deep | `DBN`, `ConvDBN`, `ResidualDBN` |
|
|
70
|
+
|
|
71
|
+
The `learnergy.core.Dataset`, `learnergy.math`, and `learnergy.visual` modules
|
|
72
|
+
remain available for array-backed datasets, SSIM/scaling helpers, convergence
|
|
73
|
+
plots, image mosaics, and tensor rendering.
|
|
74
|
+
|
|
75
|
+
See [`examples/applications`](examples/applications) for complete training and
|
|
76
|
+
classification programs.
|
|
77
|
+
|
|
78
|
+
## Development
|
|
79
|
+
|
|
80
|
+
The repository uses [uv](https://docs.astral.sh/uv/) for reproducible
|
|
81
|
+
environments and packaging:
|
|
82
|
+
|
|
83
|
+
```bash
|
|
84
|
+
uv sync --locked
|
|
85
|
+
uv run pytest
|
|
86
|
+
uv build
|
|
87
|
+
```
|
|
88
|
+
|
|
89
|
+
## Citation
|
|
90
|
+
|
|
91
|
+
```bibtex
|
|
92
|
+
@misc{roder2020learnergy,
|
|
93
|
+
title={Learnergy: Energy-based Machine Learners},
|
|
94
|
+
author={Mateus Roder and Gustavo Henrique de Rosa and João Paulo Papa},
|
|
95
|
+
year={2020},
|
|
96
|
+
eprint={2003.07443},
|
|
97
|
+
archivePrefix={arXiv},
|
|
98
|
+
primaryClass={cs.LG}
|
|
99
|
+
}
|
|
100
|
+
```
|
|
101
|
+
|
|
102
|
+
## Support
|
|
103
|
+
|
|
104
|
+
Open an [issue](https://github.com/gugarosa/learnergy/issues) for bug reports
|
|
105
|
+
and questions.
|
|
@@ -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,21 @@
|
|
|
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
|
+
originals = x.detach().cpu().numpy()
|
|
11
|
+
reconstructed = v.detach().cpu().numpy()
|
|
12
|
+
width, height = originals.shape[1:3]
|
|
13
|
+
|
|
14
|
+
return sum(
|
|
15
|
+
structural_similarity(
|
|
16
|
+
original,
|
|
17
|
+
rebuilt.reshape(width, height),
|
|
18
|
+
data_range=original.max() - original.min(),
|
|
19
|
+
)
|
|
20
|
+
for original, rebuilt in zip(originals, reconstructed)
|
|
21
|
+
) / 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
|
-
"""
|
|
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
|
+
]
|