import os.path 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 work_dir = "../../results" 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(os.path.join(work_dir, "xor_layer0_state.npz")) # 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(os.path.join(work_dir, "xor_layer0_state.npz")) # 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]")