kaiwu-torch-plugin 0.2.0__py3-none-any.whl
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.
- kaiwu/torch_plugin/__init__.py +19 -0
- kaiwu/torch_plugin/_qdiffusion_sampling.py +158 -0
- kaiwu/torch_plugin/abstract_boltzmann_machine.py +102 -0
- kaiwu/torch_plugin/dbn.py +220 -0
- kaiwu/torch_plugin/full_boltzmann_machine.py +222 -0
- kaiwu/torch_plugin/qdiffusion.py +910 -0
- kaiwu/torch_plugin/qgan.py +0 -0
- kaiwu/torch_plugin/qvae.py +180 -0
- kaiwu/torch_plugin/qvae_dist_util.py +316 -0
- kaiwu/torch_plugin/restricted_boltzmann_machine.py +148 -0
- kaiwu_torch_plugin-0.2.0.dist-info/METADATA +339 -0
- kaiwu_torch_plugin-0.2.0.dist-info/RECORD +16 -0
- kaiwu_torch_plugin-0.2.0.dist-info/WHEEL +5 -0
- kaiwu_torch_plugin-0.2.0.dist-info/licenses/LICENSE +196 -0
- kaiwu_torch_plugin-0.2.0.dist-info/licenses/NOTICE +516 -0
- kaiwu_torch_plugin-0.2.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
# -*- coding: utf-8 -*-
|
|
2
|
+
"""Kaiwu PyTorch plugin public API."""
|
|
3
|
+
|
|
4
|
+
from .dbn import UnsupervisedDBN
|
|
5
|
+
from .full_boltzmann_machine import BoltzmannMachine
|
|
6
|
+
from .qdiffusion import QDiffusion, QDiffusionConfig
|
|
7
|
+
from .qvae import QVAE
|
|
8
|
+
from .restricted_boltzmann_machine import RestrictedBoltzmannMachine
|
|
9
|
+
|
|
10
|
+
__version__ = "0.2.0"
|
|
11
|
+
|
|
12
|
+
__all__ = [
|
|
13
|
+
"RestrictedBoltzmannMachine",
|
|
14
|
+
"BoltzmannMachine",
|
|
15
|
+
"QVAE",
|
|
16
|
+
"UnsupervisedDBN",
|
|
17
|
+
"QDiffusion",
|
|
18
|
+
"QDiffusionConfig",
|
|
19
|
+
]
|
|
@@ -0,0 +1,158 @@
|
|
|
1
|
+
# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates
|
|
2
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
3
|
+
|
|
4
|
+
"""Private sampling helpers shared by the public QDiffusion model."""
|
|
5
|
+
|
|
6
|
+
from __future__ import annotations
|
|
7
|
+
|
|
8
|
+
import torch
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
# Skeptical-remasking helpers.
|
|
12
|
+
|
|
13
|
+
def topk_masking(
|
|
14
|
+
scores: torch.Tensor,
|
|
15
|
+
cutoff_len: torch.Tensor,
|
|
16
|
+
stochastic: bool = False,
|
|
17
|
+
temp: float = 1.0,
|
|
18
|
+
) -> torch.Tensor:
|
|
19
|
+
"""Selects the lowest-score positions used by skeptical remasking.
|
|
20
|
+
|
|
21
|
+
Args:
|
|
22
|
+
scores: Per-position scores used to rank editable token positions.
|
|
23
|
+
cutoff_len: Per-sample cutoff lengths that determine how many
|
|
24
|
+
positions remain masked.
|
|
25
|
+
stochastic: Whether to perturb the ranking with Gumbel noise.
|
|
26
|
+
temp: Noise temperature applied when ``stochastic`` is enabled.
|
|
27
|
+
|
|
28
|
+
Returns:
|
|
29
|
+
torch.Tensor: A boolean mask where ``True`` marks positions below the per-sample
|
|
30
|
+
cutoff.
|
|
31
|
+
"""
|
|
32
|
+
if stochastic:
|
|
33
|
+
gumbel_noise = -torch.log(-torch.log(torch.rand_like(scores) + 1e-8) + 1e-8)
|
|
34
|
+
ranked_scores = scores + temp * gumbel_noise
|
|
35
|
+
else:
|
|
36
|
+
ranked_scores = scores
|
|
37
|
+
sorted_scores = ranked_scores.sort(-1)[0]
|
|
38
|
+
cutoff = sorted_scores.gather(dim=-1, index=cutoff_len)
|
|
39
|
+
return ranked_scores < cutoff
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
# Categorical sampling helpers.
|
|
43
|
+
|
|
44
|
+
def sample_from_categorical(
|
|
45
|
+
logits: torch.Tensor,
|
|
46
|
+
temperature: float = 1.0,
|
|
47
|
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
48
|
+
"""Samples tokens from categorical logits.
|
|
49
|
+
|
|
50
|
+
Args:
|
|
51
|
+
logits: Unnormalized categorical logits.
|
|
52
|
+
temperature: Sampling temperature. A falsy value switches to greedy
|
|
53
|
+
argmax decoding.
|
|
54
|
+
|
|
55
|
+
Returns:
|
|
56
|
+
tuple[torch.Tensor, torch.Tensor]: A tuple ``(tokens, scores)``
|
|
57
|
+
where ``tokens`` contains sampled token ids and ``scores`` contains
|
|
58
|
+
the associated log-probability-style scores.
|
|
59
|
+
"""
|
|
60
|
+
if temperature:
|
|
61
|
+
dist = torch.distributions.Categorical(logits=logits.div(temperature))
|
|
62
|
+
tokens = dist.sample()
|
|
63
|
+
scores = dist.log_prob(tokens)
|
|
64
|
+
else:
|
|
65
|
+
scores, tokens = logits.log_softmax(dim=-1).max(dim=-1)
|
|
66
|
+
return tokens, scores
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def stochastic_sample_from_categorical(
|
|
70
|
+
logits: torch.Tensor,
|
|
71
|
+
temperature: float = 1.0,
|
|
72
|
+
noise_scale: float = 1.0,
|
|
73
|
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
74
|
+
"""Applies Gumbel noise before categorical sampling.
|
|
75
|
+
|
|
76
|
+
Args:
|
|
77
|
+
logits: Unnormalized categorical logits.
|
|
78
|
+
temperature: Sampling temperature forwarded to
|
|
79
|
+
:func:`sample_from_categorical`.
|
|
80
|
+
noise_scale: Multiplicative scale for the sampled Gumbel noise.
|
|
81
|
+
|
|
82
|
+
Returns:
|
|
83
|
+
tuple[torch.Tensor, torch.Tensor]: A tuple ``(tokens, scores)``
|
|
84
|
+
sampled from the perturbed categorical distribution.
|
|
85
|
+
"""
|
|
86
|
+
gumbel_noise = -torch.log(-torch.log(torch.rand_like(logits) + 1e-8) + 1e-8)
|
|
87
|
+
noisy_logits = logits + noise_scale * gumbel_noise
|
|
88
|
+
return sample_from_categorical(noisy_logits, temperature)
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def stochastic_sample_from_categorical_n(
|
|
92
|
+
logits: torch.Tensor,
|
|
93
|
+
temperature: float = 1.0,
|
|
94
|
+
noise_scale: float = 1.0,
|
|
95
|
+
n: int = 1,
|
|
96
|
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
97
|
+
"""Samples multiple noisy categorical candidates.
|
|
98
|
+
|
|
99
|
+
Args:
|
|
100
|
+
logits: Unnormalized categorical logits shaped as
|
|
101
|
+
``[batch, seq_len, vocab]`` or similar.
|
|
102
|
+
temperature: Sampling temperature forwarded to
|
|
103
|
+
:func:`sample_from_categorical`.
|
|
104
|
+
noise_scale: Multiplicative scale for the sampled Gumbel noise.
|
|
105
|
+
n: Number of independent noisy candidate sets to draw.
|
|
106
|
+
|
|
107
|
+
Returns:
|
|
108
|
+
tuple[torch.Tensor, torch.Tensor]: A tuple ``(tokens, scores)``
|
|
109
|
+
whose leading dimension indexes the sampled candidate set.
|
|
110
|
+
"""
|
|
111
|
+
expanded_logits = logits.unsqueeze(0).expand(n, *logits.shape)
|
|
112
|
+
gumbel_noise = -torch.log(
|
|
113
|
+
-torch.log(torch.rand_like(expanded_logits) + 1e-8) + 1e-8
|
|
114
|
+
)
|
|
115
|
+
noisy_logits = expanded_logits + noise_scale * gumbel_noise
|
|
116
|
+
return sample_from_categorical(noisy_logits, temperature)
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
# Logit filtering helpers.
|
|
120
|
+
|
|
121
|
+
def top_k_top_p_filtering(
|
|
122
|
+
logits: torch.Tensor,
|
|
123
|
+
top_k: int = 0,
|
|
124
|
+
top_p: float = 0.95,
|
|
125
|
+
filter_value: float = -float("Inf"),
|
|
126
|
+
) -> torch.Tensor:
|
|
127
|
+
"""Applies top-k and/or nucleus filtering to logits.
|
|
128
|
+
|
|
129
|
+
Args:
|
|
130
|
+
logits: Unnormalized categorical logits.
|
|
131
|
+
top_k: Number of highest-logit entries to retain per row. A value of
|
|
132
|
+
``0`` disables top-k filtering.
|
|
133
|
+
top_p: Nucleus-filtering threshold on cumulative probability mass.
|
|
134
|
+
filter_value: Replacement value written into filtered logit entries.
|
|
135
|
+
|
|
136
|
+
Returns:
|
|
137
|
+
torch.Tensor: A tensor with the same shape as ``logits`` where filtered entries are
|
|
138
|
+
replaced by ``filter_value``.
|
|
139
|
+
"""
|
|
140
|
+
original_shape = logits.shape
|
|
141
|
+
flat_logits = logits.reshape(-1, original_shape[-1])
|
|
142
|
+
assert flat_logits.dim() == 2
|
|
143
|
+
|
|
144
|
+
top_k = min(top_k, flat_logits.size(-1))
|
|
145
|
+
if top_k > 0:
|
|
146
|
+
indices_to_remove = (
|
|
147
|
+
flat_logits < torch.topk(flat_logits, top_k, dim=1)[0][..., -1, None]
|
|
148
|
+
)
|
|
149
|
+
flat_logits[indices_to_remove] = filter_value
|
|
150
|
+
|
|
151
|
+
sorted_logits, sorted_indices = torch.sort(flat_logits, descending=True)
|
|
152
|
+
cumulative_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
|
|
153
|
+
sorted_indices_to_remove = cumulative_probs > top_p
|
|
154
|
+
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
|
|
155
|
+
sorted_indices_to_remove[..., 0] = 0
|
|
156
|
+
sorted_logits[sorted_indices_to_remove] = filter_value
|
|
157
|
+
restored_logits = torch.gather(sorted_logits, 1, sorted_indices.argsort(-1))
|
|
158
|
+
return restored_logits.reshape(original_shape)
|
|
@@ -0,0 +1,102 @@
|
|
|
1
|
+
# -*- coding: utf-8 -*-
|
|
2
|
+
# Copyright (C) 2022-2025 Beijing QBoson Quantum Technology Co., Ltd.
|
|
3
|
+
#
|
|
4
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
"""Abstract base class for Boltzmann Machines."""
|
|
8
|
+
import torch
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class AbstractBoltzmannMachine(torch.nn.Module):
|
|
12
|
+
"""Abstract base class for Boltzmann Machines.
|
|
13
|
+
|
|
14
|
+
Args:
|
|
15
|
+
device (torch.device, optional): Device for tensor construction.
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
def __init__(self, device=None) -> None:
|
|
19
|
+
super().__init__()
|
|
20
|
+
if device is None:
|
|
21
|
+
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
22
|
+
else:
|
|
23
|
+
self.device = device
|
|
24
|
+
|
|
25
|
+
def to(self, device=..., dtype=..., non_blocking=...):
|
|
26
|
+
"""Moves the model to the specified device.
|
|
27
|
+
|
|
28
|
+
Args:
|
|
29
|
+
device: Target device.
|
|
30
|
+
dtype: Target data type.
|
|
31
|
+
non_blocking: Whether the operation should be non-blocking.
|
|
32
|
+
|
|
33
|
+
Returns:
|
|
34
|
+
AbstractBoltzmannMachine: The model on the target device.
|
|
35
|
+
"""
|
|
36
|
+
self.device = device
|
|
37
|
+
return super().to(device)
|
|
38
|
+
|
|
39
|
+
def forward(self, s_all: torch.Tensor) -> torch.Tensor:
|
|
40
|
+
"""Computes the Hamiltonian.
|
|
41
|
+
|
|
42
|
+
Args:
|
|
43
|
+
s_all (torch.Tensor): Input tensor.
|
|
44
|
+
|
|
45
|
+
Returns:
|
|
46
|
+
torch.Tensor: Hamiltonian.
|
|
47
|
+
"""
|
|
48
|
+
|
|
49
|
+
def get_ising_matrix(self):
|
|
50
|
+
"""Converts the model to Ising format.
|
|
51
|
+
|
|
52
|
+
Returns:
|
|
53
|
+
torch.Tensor: Ising matrix.
|
|
54
|
+
"""
|
|
55
|
+
return self._to_ising_matrix()
|
|
56
|
+
|
|
57
|
+
def _to_ising_matrix(self):
|
|
58
|
+
"""Converts the model to Ising format.
|
|
59
|
+
|
|
60
|
+
Returns:
|
|
61
|
+
torch.Tensor: Ising matrix.
|
|
62
|
+
|
|
63
|
+
Raises:
|
|
64
|
+
NotImplementedError: If not implemented in subclass.
|
|
65
|
+
"""
|
|
66
|
+
raise NotImplementedError("Subclasses must implement _ising method")
|
|
67
|
+
|
|
68
|
+
def objective(
|
|
69
|
+
self,
|
|
70
|
+
s_positive: torch.Tensor,
|
|
71
|
+
s_negative: torch.Tensor,
|
|
72
|
+
) -> torch.Tensor:
|
|
73
|
+
"""Objective function whose gradient is equivalent to the gradient of
|
|
74
|
+
negative log-likelihood.
|
|
75
|
+
|
|
76
|
+
Args:
|
|
77
|
+
s_positive (torch.Tensor): Tensor of observed spins (data), shape (b1, N),
|
|
78
|
+
where b1 is batch size and N is the number of variables.
|
|
79
|
+
s_negative (torch.Tensor): Tensor of spins sampled from the model, shape (b2, N),
|
|
80
|
+
where b2 is batch size and N is the number of variables.
|
|
81
|
+
|
|
82
|
+
Returns:
|
|
83
|
+
torch.Tensor: Scalar difference between data and model average energy.
|
|
84
|
+
"""
|
|
85
|
+
return self(s_positive).mean() - self(s_negative).mean()
|
|
86
|
+
|
|
87
|
+
def sample(self, sampler) -> torch.Tensor:
|
|
88
|
+
"""Samples from the Boltzmann Machine.
|
|
89
|
+
|
|
90
|
+
Args:
|
|
91
|
+
sampler (kaiwu.core.OptimizerBase): Optimizer used for sampling from the model.
|
|
92
|
+
The sampler can be kaiwuSDK's CIM or other solvers.
|
|
93
|
+
|
|
94
|
+
Returns:
|
|
95
|
+
torch.Tensor: Spins sampled from the model.
|
|
96
|
+
"""
|
|
97
|
+
ising_mat = self.get_ising_matrix()
|
|
98
|
+
solution = sampler.solve(ising_mat)
|
|
99
|
+
solution = (solution[:, :-1] * solution[:, [-1]] + 1) / 2
|
|
100
|
+
solution = torch.FloatTensor(solution)
|
|
101
|
+
solution = solution.to(self.device)
|
|
102
|
+
return solution
|
|
@@ -0,0 +1,220 @@
|
|
|
1
|
+
# -*- coding: utf-8 -*-
|
|
2
|
+
# Copyright (C) 2022-2025 Beijing QBoson Quantum Technology Co., Ltd.
|
|
3
|
+
#
|
|
4
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
5
|
+
"""Deep Belief Network (DBN) model.
|
|
6
|
+
|
|
7
|
+
This module contains the DBN class and functions for training the DBN+model
|
|
8
|
+
or only the model. Training the DBN+model will save the likelihood values
|
|
9
|
+
and prediction accuracy during the training process.
|
|
10
|
+
"""
|
|
11
|
+
import numpy as np
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
import torch
|
|
15
|
+
from torch import nn
|
|
16
|
+
|
|
17
|
+
from .restricted_boltzmann_machine import RestrictedBoltzmannMachine
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
# =================== Unsupervised DBN General Model =====================
|
|
21
|
+
class UnsupervisedDBN(nn.Module):
|
|
22
|
+
"""A general unsupervised Deep Belief Network (DBN) architecture.
|
|
23
|
+
|
|
24
|
+
This model is a stack of Restricted Boltzmann Machines (RBMs).
|
|
25
|
+
|
|
26
|
+
Args:
|
|
27
|
+
hidden_layers_structure (list, optional): A list of integers
|
|
28
|
+
representing the number of hidden units in each layer.
|
|
29
|
+
Defaults to [100, 100].
|
|
30
|
+
"""
|
|
31
|
+
|
|
32
|
+
def __init__(self, hidden_layers_structure=None):
|
|
33
|
+
super().__init__()
|
|
34
|
+
self.hidden_layers_structure = (
|
|
35
|
+
hidden_layers_structure
|
|
36
|
+
if hidden_layers_structure is not None
|
|
37
|
+
else [100, 100]
|
|
38
|
+
)
|
|
39
|
+
|
|
40
|
+
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
41
|
+
|
|
42
|
+
self.rbm_layers = None
|
|
43
|
+
self.input_dim = None
|
|
44
|
+
self._is_trained = False
|
|
45
|
+
|
|
46
|
+
def create_rbm_layer(self, input_dim):
|
|
47
|
+
"""Creates the layers of RBMs for the DBN.
|
|
48
|
+
|
|
49
|
+
Args:
|
|
50
|
+
input_dim (int): The dimension of the input data (number of visible units).
|
|
51
|
+
|
|
52
|
+
Returns:
|
|
53
|
+
UnsupervisedDBN: The instance itself with the RBM layers created.
|
|
54
|
+
"""
|
|
55
|
+
self.input_dim = input_dim
|
|
56
|
+
self.rbm_layers = nn.ModuleList()
|
|
57
|
+
|
|
58
|
+
current_dim = input_dim
|
|
59
|
+
for n_hidden in self.hidden_layers_structure:
|
|
60
|
+
rbm = RestrictedBoltzmannMachine(
|
|
61
|
+
num_visible=current_dim, # Number of visible units (feature dimension)
|
|
62
|
+
num_hidden=n_hidden, # Number of hidden units
|
|
63
|
+
).to(
|
|
64
|
+
self.device
|
|
65
|
+
) # Move model to specified device (CPU/GPU)
|
|
66
|
+
self.rbm_layers.append(rbm)
|
|
67
|
+
current_dim = n_hidden
|
|
68
|
+
|
|
69
|
+
self._is_trained = False
|
|
70
|
+
return self
|
|
71
|
+
|
|
72
|
+
def forward(self, data_in):
|
|
73
|
+
"""Performs a forward pass to transform the input data.
|
|
74
|
+
|
|
75
|
+
Args:
|
|
76
|
+
data_in (numpy.ndarray): The input data.
|
|
77
|
+
|
|
78
|
+
Returns:
|
|
79
|
+
numpy.ndarray: The transformed data after passing through all RBM layers.
|
|
80
|
+
|
|
81
|
+
Raises:
|
|
82
|
+
ValueError: If the model has not been built or trained yet.
|
|
83
|
+
"""
|
|
84
|
+
if self.rbm_layers is None:
|
|
85
|
+
raise ValueError("Model not built yet. Call create_rbm_layer first.")
|
|
86
|
+
if not self._is_trained:
|
|
87
|
+
raise ValueError(
|
|
88
|
+
"Model not trained yet. Call mark_as_trained() after training."
|
|
89
|
+
)
|
|
90
|
+
|
|
91
|
+
data_in = data_in.astype(np.float32)
|
|
92
|
+
for rbm in self.rbm_layers:
|
|
93
|
+
with torch.no_grad():
|
|
94
|
+
hidden_output = rbm.get_hidden(
|
|
95
|
+
torch.FloatTensor(data_in).to(self.device)
|
|
96
|
+
)
|
|
97
|
+
data_in = (
|
|
98
|
+
hidden_output[:, rbm.num_visible :].cpu().numpy()
|
|
99
|
+
) # Extract only the hidden part
|
|
100
|
+
return data_in
|
|
101
|
+
|
|
102
|
+
def transform(self, data_in):
|
|
103
|
+
"""An sklearn-compatible transform method.
|
|
104
|
+
|
|
105
|
+
Args:
|
|
106
|
+
data_in (numpy.ndarray): The input data.
|
|
107
|
+
|
|
108
|
+
Returns:
|
|
109
|
+
numpy.ndarray: The transformed data.
|
|
110
|
+
"""
|
|
111
|
+
return self.forward(data_in)
|
|
112
|
+
|
|
113
|
+
def reconstruct(self, data_in, layer_index=0):
|
|
114
|
+
"""Reconstructs the input from a specified RBM layer.
|
|
115
|
+
|
|
116
|
+
Args:
|
|
117
|
+
data_in (numpy.ndarray): The input data to be reconstructed.
|
|
118
|
+
|
|
119
|
+
layer_index (int, optional): The index of the RBM layer to use for reconstruction.
|
|
120
|
+
Defaults to 0.
|
|
121
|
+
|
|
122
|
+
Returns:
|
|
123
|
+
numpy.ndarray: The reconstructed data.
|
|
124
|
+
|
|
125
|
+
Raises:
|
|
126
|
+
ValueError: If the model has no RBM layers or the layer index is out of range.
|
|
127
|
+
"""
|
|
128
|
+
if self.rbm_layers is None or len(self.rbm_layers) == 0:
|
|
129
|
+
raise ValueError("No RBM layers found. Please fit the model first.")
|
|
130
|
+
|
|
131
|
+
if layer_index >= len(self.rbm_layers):
|
|
132
|
+
raise ValueError(f"Layer index {layer_index} out of range.")
|
|
133
|
+
|
|
134
|
+
rbm = self.rbm_layers[layer_index]
|
|
135
|
+
return self.reconstruct_with_rbm(rbm, data_in, self.device)
|
|
136
|
+
|
|
137
|
+
def mark_as_trained(self):
|
|
138
|
+
"""Marks the model as trained.
|
|
139
|
+
|
|
140
|
+
Returns:
|
|
141
|
+
UnsupervisedDBN: The instance itself.
|
|
142
|
+
"""
|
|
143
|
+
self._is_trained = True
|
|
144
|
+
return self
|
|
145
|
+
|
|
146
|
+
def get_rbm_layer(self, index):
|
|
147
|
+
"""Gets the RBM layer at the specified index.
|
|
148
|
+
|
|
149
|
+
Args:
|
|
150
|
+
index (int): The index of the RBM layer.
|
|
151
|
+
|
|
152
|
+
Returns:
|
|
153
|
+
RestrictedBoltzmannMachine or None: The RBM layer if found, otherwise None.
|
|
154
|
+
"""
|
|
155
|
+
if index < len(self.rbm_layers):
|
|
156
|
+
return self.rbm_layers[index]
|
|
157
|
+
return None
|
|
158
|
+
|
|
159
|
+
@staticmethod
|
|
160
|
+
def reconstruct_with_rbm(rbm, data_in, device=None):
|
|
161
|
+
"""Reconstructs data using a single RBM.
|
|
162
|
+
|
|
163
|
+
Args:
|
|
164
|
+
rbm (RestrictedBoltzmannMachine): The trained RBM model.
|
|
165
|
+
|
|
166
|
+
data_in (numpy.ndarray): The input data.
|
|
167
|
+
|
|
168
|
+
device (torch.device, optional): The device to perform computation on.
|
|
169
|
+
If None, uses the RBM's device. Defaults to None.
|
|
170
|
+
|
|
171
|
+
Returns:
|
|
172
|
+
tuple[numpy.ndarray, numpy.ndarray]: A tuple containing:
|
|
173
|
+
- The reconstructed visible layer data.
|
|
174
|
+
- The reconstruction error for each sample.
|
|
175
|
+
"""
|
|
176
|
+
if device is None:
|
|
177
|
+
device = rbm.device
|
|
178
|
+
|
|
179
|
+
# Convert to PyTorch tensor
|
|
180
|
+
data_in = torch.FloatTensor(data_in).to(device)
|
|
181
|
+
|
|
182
|
+
with torch.no_grad():
|
|
183
|
+
# Get hidden representation using RBM's get_hidden
|
|
184
|
+
hidden_act = rbm.get_hidden(data_in)
|
|
185
|
+
hidden_part = hidden_act[
|
|
186
|
+
:, rbm.num_visible :
|
|
187
|
+
] # Extract only the hidden part
|
|
188
|
+
|
|
189
|
+
# Reconstruct visible layer (using transposed weights)
|
|
190
|
+
visible_recon = torch.sigmoid(
|
|
191
|
+
torch.matmul(hidden_part, rbm.quadratic_coef.t())
|
|
192
|
+
+ rbm.linear_bias[: rbm.num_visible]
|
|
193
|
+
)
|
|
194
|
+
|
|
195
|
+
# Calculate reconstruction error
|
|
196
|
+
recon_errors = (
|
|
197
|
+
torch.mean((data_in - visible_recon) ** 2, dim=1).cpu().numpy()
|
|
198
|
+
)
|
|
199
|
+
|
|
200
|
+
return visible_recon.cpu().numpy(), recon_errors
|
|
201
|
+
|
|
202
|
+
@property
|
|
203
|
+
def num_layers(self):
|
|
204
|
+
"""Returns the number of RBM layers.
|
|
205
|
+
|
|
206
|
+
Returns:
|
|
207
|
+
int: The number of layers.
|
|
208
|
+
"""
|
|
209
|
+
return len(self.rbm_layers)
|
|
210
|
+
|
|
211
|
+
@property
|
|
212
|
+
def output_dim(self):
|
|
213
|
+
"""Returns the output dimension of the DBN.
|
|
214
|
+
|
|
215
|
+
Returns:
|
|
216
|
+
int: The dimension of the final hidden layer.
|
|
217
|
+
"""
|
|
218
|
+
if len(self.rbm_layers) > 0:
|
|
219
|
+
return self.rbm_layers[-1].num_hidden
|
|
220
|
+
return self.input_dim
|