Entity: refactored taining algo into train module
This commit is contained in:
+5
-5
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user