xor: add better documentation

This commit is contained in:
2025-12-16 18:43:54 +01:00
parent 6a3c59999e
commit 4a64279276
+13 -4
View File
@@ -109,6 +109,7 @@ class RbmLayer:
return v
def xor():
# Create params
params = RbmParams()
params.do_rao_blackwell = True
params.num_gibbs_samples = 3
@@ -117,17 +118,25 @@ def xor():
params.momentum = 0.5
params.num_epochs = 1000
status = Status()
# Create layer
layer = RbmLayer("Layer_0", 3, 16, params)
# Init weights
layer.init(0.01)
# Load weights (if exists)
layer.load()
# Train
# Prepare training data
training_batch = np.array([[0,1,1], [0,0,0], [1,1,0], [1,0,1]], dtype=np.float64)
layer.train(training_batch, cd_jens, status)
# Train layer
layer.train(training_batch, cd_jens, Status())
# Save weights
layer.save()
# Test
# Test with test data
test_batch = np.array([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64)
for pattern in test_batch:
h = layer.gibbs_v_to_h(pattern)