- numGibbs no longer part of EntityParameter

- refactored forward and reconstruct
- conditionally use optimizer for training
This commit is contained in:
2025-12-19 17:15:10 +01:00
parent a47922cb1c
commit d3c9fe4681
4 changed files with 124 additions and 41 deletions
+11 -5
View File
@@ -2,11 +2,13 @@ import os.path
from rbm.layer import Layer
from rbm.status import Status
from rbm.train import train, TrainingParams
from rbm.train import train, TrainingParams, Optimizer
from rbm.matrix import Mat, np
from rbm.entity import EntityParams
work_dir = "../../results"
WORK_DIR = "../../results"
USE_OPTIMIZER = True
def xor():
# Create params
entity_params = EntityParams()
@@ -21,16 +23,20 @@ def xor():
layer.init(0.01)
# Load weights (if exists)
layer.load(os.path.join(work_dir, "xor_layer0_state.npz"))
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)
# Train layer
train(layer.entity, training_batch, training_params, Status())
if USE_OPTIMIZER:
optim = Optimizer(layer.entity, training_params)
optim(training_batch, Status())
else:
train(layer.entity, training_batch, training_params, Status())
# Save weights
layer.save(os.path.join(work_dir, "xor_layer0_state.npz"))
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)