Entity: refactored taining algo into train module

This commit is contained in:
2025-12-18 21:23:16 +01:00
parent 6664ebdb9b
commit b206ab5788
7 changed files with 120 additions and 118 deletions
+5 -5
View File
@@ -1,12 +1,12 @@
from params import RbmParams
from params import EntityParams
from state import RbmState
from status import Status
from cd_train import cd_jens
from train import train
from entity import Entity
from matrix import Mat, np
class Layer:
def __init__(self, name: str, shape: tuple[int, int, int, int], params: RbmParams):
def __init__(self, name: str, shape: tuple[int, int, int, int], params: EntityParams):
self.name = name
self.shape = shape
self.entity = Entity((shape[0]*shape[1]+shape[2], shape[3]), params)
@@ -29,7 +29,7 @@ class Layer:
def xor():
# Create params
params = RbmParams()
params = EntityParams()
params.do_rao_blackwell = True
params.num_gibbs_samples = 3
@@ -46,7 +46,7 @@ def xor():
training_batch = Mat([[0,1,1], [0,0,0], [1,1,0], [1,0,1]], dtype=np.float64)
# Train layer
layer.entity.train(training_batch, cd_jens, Status())
train(layer.entity, training_batch, Status())
# Save weights
layer.save()