diff --git a/src/rbm/Layer_0_state.npz b/src/rbm/Layer_0_state.npz deleted file mode 100644 index ae6146b..0000000 Binary files a/src/rbm/Layer_0_state.npz and /dev/null differ diff --git a/src/rbm/layer.py b/src/rbm/layer.py index b11ec15..7f65f97 100644 --- a/src/rbm/layer.py +++ b/src/rbm/layer.py @@ -10,19 +10,14 @@ class Layer: self.name = name self.shape = shape self.entity = Entity((shape[0]*shape[1]+shape[2], shape[3]), params) - self.state_filename = f"{self.name}_state.npz" def init(self, std: float): self.entity.state.init(mu=0, std=std) def save(self, filename: str = None): - if filename is None: - filename = self.state_filename self.entity.state.to_file(filename) def load(self, filename: str = None): - if filename is None: - filename = self.state_filename state = RbmState.from_file(filename) if state is not None: self.entity.state = state diff --git a/src/rbm/status.py b/src/rbm/status.py index b0c4c5e..b3108b2 100644 --- a/src/rbm/status.py +++ b/src/rbm/status.py @@ -29,4 +29,5 @@ class Status: return False def on_report(self, status: dict) -> bool: + Status.print_status(status) return True \ No newline at end of file diff --git a/src/rbm/test_rbm.py b/src/tests/test_rbm.py similarity index 94% rename from src/rbm/test_rbm.py rename to src/tests/test_rbm.py index 95f11db..133de36 100644 --- a/src/rbm/test_rbm.py +++ b/src/tests/test_rbm.py @@ -1,10 +1,10 @@ import os.path import cv2 as cv from argparse import ArgumentParser -from .stack_factory import StackFactory -from .status import Status -from .stack_deep import StackDeep -from .matrix import Mat, np, convert +from rbm.stack_factory import StackFactory +from rbm.status import Status +from rbm.stack_deep import StackDeep +from rbm.matrix import Mat, np, convert def cv_show(name: str, vec: Mat, shape): img = cv.Mat(convert(np.resize(vec, shape))) diff --git a/src/rbm/test_xor.py b/src/tests/test_xor.py similarity index 84% rename from src/rbm/test_xor.py rename to src/tests/test_xor.py index 50a9a76..248aa68 100644 --- a/src/rbm/test_xor.py +++ b/src/tests/test_xor.py @@ -1,9 +1,12 @@ +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() @@ -17,7 +20,7 @@ def xor(): layer.init(0.01) # Load weights (if exists) - layer.load() + 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) @@ -26,7 +29,7 @@ def xor(): train(layer.entity, training_batch, Status()) # Save weights - layer.save() + 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)