From a9e67da43a160eff1adba5637b53d5bc3c63b09c Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Fri, 19 Dec 2025 11:55:29 +0100 Subject: [PATCH] refactored --- src/rbm/Layer_0_state.npz | Bin 1268 -> 0 bytes src/rbm/layer.py | 5 ----- src/rbm/status.py | 1 + src/{rbm => tests}/test_rbm.py | 8 ++++---- src/{rbm => tests}/test_xor.py | 7 +++++-- 5 files changed, 10 insertions(+), 11 deletions(-) delete mode 100644 src/rbm/Layer_0_state.npz rename src/{rbm => tests}/test_rbm.py (94%) rename src/{rbm => tests}/test_xor.py (84%) diff --git a/src/rbm/Layer_0_state.npz b/src/rbm/Layer_0_state.npz deleted file mode 100644 index ae6146bbbdf232a9a30067c938aaeb7c7706f76a..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 1268 zcmWIWW@gc4fB;1X%?(+>|Dk}LL4=_^qf9Tappub6fPsMtstQU^_6zk5h-73aW2jb7 zNi9w;Qnyl2w@I^5*HKVU%P%S^O3aJTFG@)TiMu7{6sH2ki!%}nQh|I8V;u!UGff?Z zS_N_e*R+BvQ~!dm`?H%;m46)PI$)WT_i3re6?^}kji1Zj{<9YdblGuT{-6CaWf8+u zo4)Kn-xqq|;uH1*ADtfTexby1U1)up=O^K0y40{T?0rQ}6V@ z*e8nRawytyAK;1jzu^7+@Aj9{R(>`MdTeie_|oUc3qIP{>6t3=vfj04-u0M|*@NwX ziRXnmHq(FEN8NN{Jo=C2z@PNGBS&ln4ovJ-+|rp`z|jk;&G6S=DV4v%;F$6 zInIke=&F5%H$#BJ?Em}C^DB#=d|-3v5fd{=`Fno$y=fH!7sn5EF85}%!uIO_=$95p#MqI|yqYMr&6{jXN1hE{@PYaf5 z`^w}n%l!TR&r!SWzbU*-oV99){fWBd54YmK+3)!>Z`NFwxAvP49}Ej#{eFJ{O2YBw ztJQ`j95!IWNrI=F8PH?{qbW)^hB^wy6eSx4El|9Qtz03F55e5zfV4S@~eG${HmtqlaB4r_gJqu)BV3aPseN(%Y9$=KQ3ZF zRLc0vJ}={0<{IfW_UqYy8=U#`+CFwdFO#wTo&C`P!Eqj|zwZC?up$1#<0tmhcJ)^- zX}D@1;LXUS%ZyrtfXfPCIl&GMDiGBG4lihdg02aad_V~VgxNr(~` 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)