diff --git a/src/rbm/entity.py b/src/rbm/entity.py index 65a6c85..734f37e 100644 --- a/src/rbm/entity.py +++ b/src/rbm/entity.py @@ -32,17 +32,8 @@ class TrainingParams: @classmethod def from_dict(cls, params: dict): obj = TrainingParams() - obj.learning_rate = params["learningRate"] - obj.momentum = params["momentum"] - obj.weight_decay = params["weightDecay"] - obj.l2_lambda = params["l2_lambda"] - obj.num_epochs = params["numEpochs"] - obj.num_gibbs_samples = params["numGibbs"] - obj.mini_batch_size = params["miniBatchSize"] - obj.do_rao_blackwell = params["doRaoBlackwell"] - obj.do_gibbs_sample_visible = params["gibbsDoSampleVisible"] - obj.do_gibbs_sample_hidden = params["gibbsDoSampleHidden"] - obj.do_batch_sample = params["doSampleBatch"] + for key, value in params.items(): + setattr(obj, key, value) return obj diff --git a/src/tests/test_rbm.py b/src/tests/test_rbm.py index d6bd2bf..37643f5 100644 --- a/src/tests/test_rbm.py +++ b/src/tests/test_rbm.py @@ -5,6 +5,7 @@ from rbm.stack_factory import StackFactory from rbm.status import Status from rbm.stack_deep import StackDeep from rbm.matrix import Mat, np, convert, read_armadillo +from rbm.entity import Entity def cv_show(name: str, vec: Mat, shape): img = cv.Mat(convert(np.resize(vec, shape))) @@ -18,10 +19,10 @@ class MyStatus(Status): self.batch = _batch self.index = 0 - def on_report(self, status: dict) -> bool: + def on_report(self, entity: Entity, status: dict) -> bool: do_continue = True # print status values - Status.print_status(status) + Status.print_status(entity, status) # Shape of training vector shape = self.stack.from_index(0).shape[0:2] + (1,)