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.
- learnergy-2.0.1/PKG-INFO +168 -0
- learnergy-2.0.1/README.md +131 -0
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/__init__.py +1 -1
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/core/__init__.py +3 -2
- learnergy-2.0.1/learnergy/core/dataset.py +48 -0
- learnergy-2.0.1/learnergy/core/model.py +58 -0
- learnergy-2.0.1/learnergy/math/__init__.py +1 -0
- learnergy-2.0.1/learnergy/math/metrics.py +30 -0
- learnergy-2.0.1/learnergy/math/scale.py +14 -0
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/__init__.py +1 -2
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/bernoulli/__init__.py +12 -3
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/bernoulli/conv_rbm.py +71 -224
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/bernoulli/discriminative_rbm.py +39 -127
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/bernoulli/dropout_rbm.py +16 -104
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/bernoulli/e_dropout_rbm.py +24 -50
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/bernoulli/rbm.py +47 -191
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/deep/__init__.py +3 -2
- learnergy-2.0.1/learnergy/models/deep/conv_dbn.py +247 -0
- learnergy-2.0.1/learnergy/models/deep/dbn.py +231 -0
- learnergy-2.0.1/learnergy/models/deep/residual_dbn.py +130 -0
- learnergy-2.0.1/learnergy/models/extra/__init__.py +5 -0
- learnergy-2.0.1/learnergy/models/extra/sigmoid_rbm.py +34 -0
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/gaussian/__init__.py +12 -2
- learnergy-2.0.1/learnergy/models/gaussian/_normalization.py +13 -0
- learnergy-2.0.1/learnergy/models/gaussian/gaussian_conv_rbm.py +128 -0
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy/models/gaussian/gaussian_rbm.py +45 -331
- learnergy-2.0.1/learnergy/utils/__init__.py +1 -0
- learnergy-2.0.1/learnergy/utils/constants.py +3 -0
- learnergy-2.0.1/learnergy/utils/exception.py +40 -0
- learnergy-2.0.1/learnergy/utils/logging.py +37 -0
- learnergy-2.0.1/learnergy/visual/__init__.py +1 -0
- learnergy-2.0.1/learnergy/visual/convergence.py +41 -0
- learnergy-2.0.1/learnergy/visual/image.py +97 -0
- learnergy-2.0.1/learnergy/visual/tensor.py +38 -0
- learnergy-2.0.1/learnergy.egg-info/PKG-INFO +168 -0
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy.egg-info/SOURCES.txt +6 -3
- learnergy-2.0.1/learnergy.egg-info/requires.txt +18 -0
- learnergy-2.0.1/pyproject.toml +71 -0
- learnergy-2.0.1/setup.cfg +4 -0
- learnergy-2.0.1/tests/test_bernoulli.py +234 -0
- learnergy-2.0.1/tests/test_core.py +23 -0
- learnergy-2.0.1/tests/test_deep.py +207 -0
- learnergy-2.0.1/tests/test_gaussian.py +226 -0
- learnergy-2.0.1/tests/test_utilities.py +127 -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.1}/LICENSE +0 -0
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy.egg-info/dependency_links.txt +0 -0
- {learnergy-1.2.0 → learnergy-2.0.1}/learnergy.egg-info/top_level.txt +0 -0
learnergy-2.0.1/PKG-INFO
ADDED
|
@@ -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
|
+
[](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. 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
|
+
[](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. 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.
|
|
@@ -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
|
-
"""
|
|
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
|
+
]
|