- removed Optimizer

- refactored training
- use static seed for random (for better comparison)
- simplified StackDeep.train
This commit is contained in:
2025-12-20 20:12:19 +01:00
parent 5f5c7c6d77
commit 61c762150c
5 changed files with 30 additions and 116 deletions
+2 -6
View File
@@ -2,7 +2,7 @@ import os.path
from rbm.layer import Layer
from rbm.status import Status
from rbm.train import train, TrainingParams, Optimizer
from rbm.train import train, TrainingParams
from rbm.matrix import Mat, np
from rbm.entity import EntityParams
@@ -29,11 +29,7 @@ def xor():
training_batch = Mat([[0,1,1], [0,0,0], [1,1,0], [1,0,1]], dtype=np.float64)
# Train layer
if USE_OPTIMIZER:
optim = Optimizer(layer.entity, training_params)
optim(training_batch, Status())
else:
train(layer.entity, training_batch, training_params, Status())
train(layer.entity, training_batch, training_params, Status())
# Save weights
layer.save(os.path.join(WORK_DIR, "xor_layer0_state.npz"))