Gaussian-Bernoulli Restricted Boltzmann Machines (GRBM)
| 5 | from utils import cosine_schedule |
| 6 | |
| 7 | class 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] |