From 13a7cff094818f28a483727493162b39bdcd7295 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Fri, 19 Dec 2025 11:39:55 +0100 Subject: [PATCH] layer: refactored xor() into test_xor.py --- src/rbm/Layer_0_state.npz | Bin 0 -> 1268 bytes src/rbm/layer.py | 36 ---------------------------------- src/rbm/test_xor.py | 40 ++++++++++++++++++++++++++++++++++++++ 3 files changed, 40 insertions(+), 36 deletions(-) create mode 100644 src/rbm/Layer_0_state.npz create mode 100644 src/rbm/test_xor.py diff --git a/src/rbm/Layer_0_state.npz b/src/rbm/Layer_0_state.npz new file mode 100644 index 0000000000000000000000000000000000000000..ae6146bbbdf232a9a30067c938aaeb7c7706f76a GIT binary patch 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(~` literal 0 HcmV?d00001 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]")