MCPcopy Create free account
hub / github.com/DSL-Lab/GRBM / GRBM

Class GRBM

grbm.py:7–361  ·  view source on GitHub ↗

Gaussian-Bernoulli Restricted Boltzmann Machines (GRBM)

Source from the content-addressed store, hash-verified

5from utils import cosine_schedule
6
7class GRBM(nn.Module):
8 """ Gaussian-Bernoulli Restricted Boltzmann Machines (GRBM) """
9
10 def __init__(self,
11 visible_size,
12 hidden_size,
13 CD_step=1,
14 CD_burnin=0,
15 init_var=1e-0,
16 inference_method='Gibbs',
17 Langevin_step=10,
18 Langevin_eta=1.0,
19 is_anneal_Langevin=True,
20 Langevin_adjust_step=0) -> None:
21 super().__init__()
22 # we use samples in [CD_burnin, CD_step) steps
23 assert CD_burnin >= 0 and CD_burnin <= CD_step
24 assert inference_method in ['Gibbs', 'Langevin', 'Gibbs-Langevin']
25
26 self.visible_size = visible_size
27 self.hidden_size = hidden_size
28 self.CD_step = CD_step
29 self.CD_burnin = CD_burnin
30 self.init_var = init_var
31 self.inference_method = inference_method
32 self.Langevin_step = Langevin_step
33 self.Langevin_eta = Langevin_eta
34 self.is_anneal_Langevin = is_anneal_Langevin
35 self.Langevin_adjust_step = Langevin_adjust_step
36
37 self.W = nn.Parameter(torch.Tensor(visible_size, hidden_size))
38 self.b = nn.Parameter(torch.Tensor(hidden_size))
39 self.mu = nn.Parameter(torch.Tensor(visible_size))
40 self.log_var = nn.Parameter(torch.Tensor(visible_size))
41 self.reset_parameters()
42
43 def reset_parameters(self):
44 nn.init.normal_(self.W,
45 std=1.0 * self.init_var /
46 np.sqrt(self.visible_size + self.hidden_size))
47 nn.init.constant_(self.b, 0.0)
48 nn.init.constant_(self.mu, 0.0)
49 nn.init.constant_(self.log_var,
50 np.log(self.init_var)) # init variance = 1.0
51
52 def get_var(self):
53 return self.log_var.exp().clip(min=1e-8)
54
55 def set_Langevin_eta(self, eta):
56 self.Langevin_eta = eta
57
58 def set_Langevin_adjust_step(self, step):
59 self.Langevin_adjust_step = step
60
61 @torch.no_grad()
62 def energy(self, v, h):
63 # compute per-sample energy averaged over batch size
64 B = v.shape[0]

Callers 1

train_modelFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected