diff --git a/src/tests/test_linear.py b/src/tests/test_linear.py new file mode 100644 index 0000000..22fd60f --- /dev/null +++ b/src/tests/test_linear.py @@ -0,0 +1,46 @@ +import os.path + +from rbm.layer import Layer +from rbm.status import Status +from rbm.train import train +from rbm.matrix import Mat, np +from rbm.entity import EntityParams, TrainingParams + +WORK_DIR = "../../results" +USE_OPTIMIZER = True + +def linear(): + # Create params + entity_params = EntityParams(do_gaussian_visible=False, do_gaussian_hidden=False) + training_params = TrainingParams(learning_rate=0.0001, do_batch_sample=True, num_epochs=10000, momentum=0.9) + entity_params.do_rao_blackwell = False + entity_params.num_gibbs_samples = 1 + + # Create layer + layer = Layer("Layer_0", (3, 1, 0, 32), entity_params, training_params) + + # Init weights + layer.init(0.01) + + # Load weights (if exists) + layer.load(os.path.join(WORK_DIR, "linear_layer0_state.npz")) + + # Prepare training data + training_batch = Mat([[0.5,0.5,1], [0.1,0.9,1.0], [0.2,0.5,0.7], [0.9,0.1,1], [0.5,0.2,0.7]], dtype=np.float64) + + # Train layer + train(layer.entity, training_batch, Status()) + + # Save weights + layer.save(os.path.join(WORK_DIR, "linear_layer0_state.npz")) + + # Test with test data + test_batch = training_batch + for pattern in test_batch: + h = layer.entity.forward(pattern) + v = layer.entity.reconstruct(h) + print(f"P{pattern} : {v}") + +if __name__ == "__main__": + linear() + print("Test: [passed]") diff --git a/src/tests/test_xor.py b/src/tests/test_xor.py index 8a8adc2..2ed35bb 100644 --- a/src/tests/test_xor.py +++ b/src/tests/test_xor.py @@ -2,9 +2,9 @@ import os.path from rbm.layer import Layer from rbm.status import Status -from rbm.train import train, TrainingParams +from rbm.train import train from rbm.matrix import Mat, np -from rbm.entity import EntityParams +from rbm.entity import EntityParams, TrainingParams WORK_DIR = "../../results" USE_OPTIMIZER = True @@ -29,7 +29,7 @@ def xor(): 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, training_params, Status()) + train(layer.entity, training_batch, Status()) # Save weights layer.save(os.path.join(WORK_DIR, "xor_layer0_state.npz"))