diff --git a/src/rbm/entity.py b/src/rbm/entity.py index df53acf..2f75592 100644 --- a/src/rbm/entity.py +++ b/src/rbm/entity.py @@ -2,10 +2,10 @@ from .state import RbmState from .matrix import prob, Mat class EntityParams: - def __init__(self): + def __init__(self, do_gaussian_visible: bool = False, do_gaussian_hidden: bool = False): # Entity parameters - self.do_gaussian_visible = False - self.do_gaussian_hidden = False + self.do_gaussian_visible = do_gaussian_visible + self.do_gaussian_hidden = do_gaussian_hidden @classmethod def from_dict(cls, params: dict): @@ -24,6 +24,9 @@ class Entity: self.state = RbmState.from_layer_params(shape) self.grad = RbmState.from_layer_params(shape) + def __call__(self, x: Mat): + return self.forward(x) + def grad_zero(self): self.grad = RbmState.from_layer_params(self.shape) diff --git a/src/rbm/model.py b/src/rbm/model.py new file mode 100644 index 0000000..ae5b333 --- /dev/null +++ b/src/rbm/model.py @@ -0,0 +1,45 @@ +import os +from abc import ABC, abstractmethod + +from .entity import Entity +from .train import train, TrainingParams, cd_jens +from .matrix import Mat, np +from .status import Status +from .state import RbmState + +class Model(ABC): + known_classes = [Entity] + def __init__(self, name: str = "myStack", work_dir: str = "."): + self.name = name + self.work_dir = work_dir + + def objects(self, obj_type: type = Entity) -> list[Entity]: + obj_list: list[type[obj_type]] = [] + for name, value in self.__dict__.items(): + if isinstance(value, obj_type): + obj_list.append(value) + return obj_list + + def train(self, batch: Mat, params: TrainingParams): + _batch = np.copy(batch) + for entity in self.objects(Entity): + train(entity, _batch, params, Status(), cd_jens) + _batch = entity(_batch) + + @abstractmethod + def forward(self, x: Mat) -> Mat: + pass + + def save(self): + for index, entity in enumerate(self.objects(Entity)): + filepath = os.path.join(self.work_dir, f"{self.name}-{index}-state.npz") + entity.state.to_file(filepath) + + def load(self): + for index, entity in enumerate(self.objects(Entity)): + filepath = os.path.join(self.work_dir, f"{self.name}-{index}-state.npz") + state = RbmState.from_file(filepath) + if state is not None: + entity.state = state + + diff --git a/src/rbm/train.py b/src/rbm/train.py index 23e3942..a6e5d54 100644 --- a/src/rbm/train.py +++ b/src/rbm/train.py @@ -4,18 +4,29 @@ from .entity import Entity from .status import Status class TrainingParams: - def __init__(self): + 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 = 0.1 - self.momentum = 0.5 - self.weight_decay = 0 - self.num_epochs = 1000 - self.mini_batch_size = 0 - self.do_rao_blackwell = False - self.do_gibbs_sample_visible = False - self.do_gibbs_sample_hidden = False - self.do_batch_sample = False - self.num_gibbs_samples = 1 + 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): @@ -24,12 +35,12 @@ class TrainingParams: 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"] - obj.num_gibbs_samples = params["numGibbs"] return obj @@ -70,7 +81,7 @@ def cd_jens(entity: Entity, v_states: Mat, params: TrainingParams): dbv -= np.sum(v_probs, 0) dbh -= np.sum(h_probs, 0) - return dw, dbv, dbh, h_probs + return dw, dbv, dbh def to_mini_batch(batch: Mat, mini_batch_size: int): mini_batches = [] @@ -97,7 +108,7 @@ def train(entity: Entity, batch: Mat, params: TrainingParams, status: Status, cd entity.grad_zero() for epochs in range(params.num_epochs): # Contrastive divergence learning: calculate gradients - dwhv, dbv, dbh, _ = cd_func(entity, mini_batch, params) + dwhv, dbv, dbh = cd_func(entity, mini_batch, params) # Adjust weight and biases grad = entity.grad_compute(dbv, dbh, dwhv, learning_rate=params.learning_rate/batch.shape[0], momentum=params.momentum, weight_decay=params.weight_decay) diff --git a/src/tests/test_model.py b/src/tests/test_model.py new file mode 100644 index 0000000..23e06eb --- /dev/null +++ b/src/tests/test_model.py @@ -0,0 +1,40 @@ +from sympy.codegen.ast import float64 + +from rbm.model import Model +from rbm.entity import Entity, EntityParams +from rbm.matrix import Mat, np +from rbm.train import TrainingParams + +class DeepStack(Model): + def __init__(self, name: str = "myStack"): + super().__init__(name) + self.unit1 = Entity((32*32*3, 256), EntityParams(do_gaussian_visible=True)) +# self.unit2 = Entity((16, 24), EntityParams()) +# self.unit3 = Entity((24, 10), EntityParams()) + + def forward(self, x: Mat): + x = self.unit1(x) +# x = self.unit2(x) +# x = self.unit3(x) + return x + + +if __name__ == "__main__": + # Create model + model = DeepStack("myStack") + + # load state + model.load() + + # create batch + batch = (np.random.rand(16, 32*32*3) > 0.5).astype(np.float64) + + # Train + model.train(batch, TrainingParams(learning_rate=0.1, momentum=0.9, do_rao_blackwell=True, num_epochs=1000)) + + # save state + model.save() + + for pat in batch: + model.forward(pat) +