diff --git a/src/rbm/Layer_0_state.npz b/src/rbm/Layer_0_state.npz new file mode 100644 index 0000000..ae6146b Binary files /dev/null and b/src/rbm/Layer_0_state.npz differ diff --git a/src/rbm/layer.py b/src/rbm/layer.py index 94b4535..b11ec15 100644 --- a/src/rbm/layer.py +++ b/src/rbm/layer.py @@ -27,42 +27,6 @@ class Layer: if state is not None: self.entity.state = state -def xor(): - # Create params - params = EntityParams() - params.do_rao_blackwell = True - params.num_gibbs_samples = 3 - - # Create layer - layer = Layer("Layer_0", (3, 1, 0, 16), params) - - # Init weights - layer.init(0.01) - - # Load weights (if exists) - layer.load() - - # Prepare training data - training_batch = Mat([[0,1,1], [0,0,0], [1,1,0], [1,0,1]], dtype=np.float64) - - # Train layer - train(layer.entity, training_batch, Status()) - - # Save weights - layer.save() - - # Test with test data - test_batch = Mat([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64) - for pattern in test_batch: - h = layer.entity.gibbs_v_to_h(pattern) - v = layer.entity.gibbs_h_to_v(h) - print(f"P{pattern} : {v}") - - -if __name__ == "__main__": - xor() - print("Test: [passed]") - diff --git a/src/rbm/test_xor.py b/src/rbm/test_xor.py new file mode 100644 index 0000000..50a9a76 --- /dev/null +++ b/src/rbm/test_xor.py @@ -0,0 +1,40 @@ +from rbm.params import EntityParams +from rbm.layer import Layer +from rbm.status import Status +from rbm.train import train +from rbm.matrix import Mat, np + +def xor(): + # Create params + params = EntityParams() + params.do_rao_blackwell = True + params.num_gibbs_samples = 3 + + # Create layer + layer = Layer("Layer_0", (3, 1, 0, 16), params) + + # Init weights + layer.init(0.01) + + # Load weights (if exists) + layer.load() + + # Prepare training data + training_batch = Mat([[0,1,1], [0,0,0], [1,1,0], [1,0,1]], dtype=np.float64) + + # Train layer + train(layer.entity, training_batch, Status()) + + # Save weights + layer.save() + + # Test with test data + test_batch = Mat([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64) + for pattern in test_batch: + h = layer.entity.gibbs_v_to_h(pattern) + v = layer.entity.gibbs_h_to_v(h) + print(f"P{pattern} : {v}") + +if __name__ == "__main__": + xor() + print("Test: [passed]")