- training params are (again) attrubute of Entity
- ditched layout concept
This commit is contained in:
+44
-1
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user