refactored

This commit is contained in:
2025-12-19 15:20:04 +01:00
parent a293fc31a0
commit a47922cb1c
8 changed files with 92 additions and 124 deletions
+10 -9
View File
@@ -1,20 +1,21 @@
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.train import train, TrainingParams
from rbm.matrix import Mat, np
from rbm.entity import EntityParams
work_dir = "../../results"
def xor():
# Create params
params = EntityParams()
params.do_rao_blackwell = True
params.num_gibbs_samples = 3
entity_params = EntityParams()
training_params = TrainingParams()
entity_params.do_rao_blackwell = True
entity_params.num_gibbs_samples = 3
# Create layer
layer = Layer("Layer_0", (3, 1, 0, 16), params)
layer = Layer("Layer_0", (3, 1, 0, 16), entity_params, training_params)
# Init weights
layer.init(0.01)
@@ -26,7 +27,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, Status())
train(layer.entity, training_batch, training_params, Status())
# Save weights
layer.save(os.path.join(work_dir, "xor_layer0_state.npz"))
@@ -34,8 +35,8 @@ def xor():
# 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)
h = layer.entity.forward(pattern)
v = layer.entity.reconstruct(h)
print(f"P{pattern} : {v}")
if __name__ == "__main__":