diff --git a/src/rbm/layer.py b/src/rbm/layer.py index 38300ad..61c72fa 100644 --- a/src/rbm/layer.py +++ b/src/rbm/layer.py @@ -13,6 +13,9 @@ class RbmLayer: self.params = params self.state_filename = f"{self.name}_state.npz" + def init(self): + self.state.init() + def save(self): self.state.to_file(self.state_filename) @@ -89,6 +92,22 @@ class RbmLayer: return prob(state) + def gibbs_v_to_h(self, v: np.ndarray) -> np.ndarray: + h = None + for i in range(self.params.num_gibbs_samples): + h = self.v_to_ph(v) + v = self.h_to_pv(h) + + return h + + def gibbs_h_to_v(self, h: np.ndarray) -> np.ndarray: + v = None + for i in range(self.params.num_gibbs_samples): + v = self.h_to_pv(h) + h = self.v_to_ph(v) + + return v + def xor(): params = RbmParams() params.do_rao_blackwell = True @@ -96,15 +115,24 @@ def xor(): params.mini_batch_size = 100 params.learning_rate = 0.1 params.momentum = 0.5 - params.num_epochs = 10000 + params.num_epochs = 100 status = Status() layer = RbmLayer("Layer_0", 3, 16, params) - training_batch = np.array([[0,0,0], [0,1,1], [1,0,1], [1,1,0]], dtype=np.float64) - test_batch = np.array([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64) + layer.init() + # Train + training_batch = np.array([[0,0,0], [0,1,1], [1,0,1], [1,1,0]], dtype=np.float64) layer.train(training_batch, cd_jens, status) + # Test + test_batch = np.array([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64) + for pattern in test_batch: + h = layer.gibbs_v_to_h(pattern) + v = layer.gibbs_h_to_v(h) + print(f"P{pattern} : {v}") + + if __name__ == "__main__": xor()