- training params are (again) attrubute of Entity

- ditched layout concept
This commit is contained in:
2026-01-02 14:51:19 +01:00
parent 8b97350564
commit 10691894d1
8 changed files with 70 additions and 98 deletions
+44 -1
View File
@@ -1,6 +1,48 @@
from .state import RbmState
from .matrix import prob, Mat
from enum import Enum
class TrainingParams:
def __init__(self,
learning_rate: float = 0.1,
momentum: float = 0.5,
weight_decay: float = 0.0,
num_epochs: int = 1000,
num_gibbs_samples: int = 1,
mini_batch_size: int = 0,
do_rao_blackwell: bool = False,
do_gibbs_sample_visible: bool = False,
do_gibbs_sample_hidden: bool = False,
do_batch_sample: bool = False
):
# Training parameters
self.learning_rate = learning_rate
self.momentum = momentum
self.weight_decay = weight_decay
self.num_epochs = num_epochs
self.num_gibbs_samples = num_gibbs_samples
self.mini_batch_size = mini_batch_size
self.do_rao_blackwell = do_rao_blackwell
self.do_gibbs_sample_visible = do_gibbs_sample_visible
self.do_gibbs_sample_hidden = do_gibbs_sample_hidden
self.do_batch_sample = do_batch_sample
@classmethod
def from_dict(cls, params: dict):
obj = TrainingParams()
obj.learning_rate = params["learningRate"]
obj.momentum = params["momentum"]
obj.weight_decay = params["weightDecay"]
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"]
return obj
class EntityParams:
def __init__(self, do_gaussian_visible: bool = False, do_gaussian_hidden: bool = False, num_gibbs_samples: int = 1):
# Entity parameters
@@ -27,9 +69,10 @@ class Entity:
GB_RBM = "GB-RBM"
GG_RBM = "GG-RBM"
def __init__(self, shape: tuple[int, int], params: EntityParams):
def __init__(self, shape: tuple[int, int], params: EntityParams, training_params: TrainingParams|None = None):
self.shape = shape
self.params = params
self.training_params = training_params
self.state = RbmState.from_layer_params(shape)
self.grad = RbmState.from_layer_params(shape)
self.type = None