- state load/save revised

- num gibbs samples is also part of entity
- working 3-layer deep test model
This commit is contained in:
2025-12-21 17:34:11 +01:00
parent 22d189cf2f
commit 4bd5cffd6b
5 changed files with 57 additions and 30 deletions
+24 -14
View File
@@ -5,36 +5,46 @@ from rbm.entity import Entity, EntityParams
from rbm.matrix import Mat, np
from rbm.train import TrainingParams
class DeepStack(Model):
def __init__(self, name: str = "myStack"):
super().__init__(name)
self.unit1 = Entity((32*32*3, 256), EntityParams(do_gaussian_visible=True))
# self.unit2 = Entity((16, 24), EntityParams())
# self.unit3 = Entity((24, 10), EntityParams())
class TestModel(Model):
def __init__(self, name: str, work_dir: str = '.'):
super().__init__(name, work_dir)
self.unit1 = Entity((16, 64), EntityParams())
self.unit2 = Entity((64, 16), EntityParams())
self.unit3 = Entity((16, 64), EntityParams())
def forward(self, x: Mat):
x = self.unit1(x)
# x = self.unit2(x)
# x = self.unit3(x)
x = self.unit1.forward(x)
x = self.unit2.forward(x)
x = self.unit3.forward(x)
return x
def backward(self, x: Mat):
x = self.unit3.reconstruct(x)
x = self.unit2.reconstruct(x)
x = self.unit1.reconstruct(x)
return x
if __name__ == "__main__":
# Create model
model = DeepStack("myStack")
model = TestModel("TestModel", "results")
# Init state
model.init(0.01)
# load state
model.load()
# create batch
batch = (np.random.rand(16, 32*32*3) > 0.5).astype(np.float64)
batch = (np.random.rand(64, 16) > 0.5).astype(np.float64)
# Train
model.train(batch, TrainingParams(learning_rate=0.1, momentum=0.9, do_rao_blackwell=True, num_epochs=1000))
model.train(batch, TrainingParams(learning_rate=0.01, momentum=0.9, do_rao_blackwell=True, num_epochs=1000, num_gibbs_samples=3))
# save state
model.save()
for pat in batch:
model.forward(pat)
for inp in batch:
out = model.backward(model.forward(inp))
print(f"Pattern: {inp} -> {out}")