added model

This commit is contained in:
2025-12-21 16:39:53 +01:00
parent 7668c73ea4
commit 22d189cf2f
4 changed files with 116 additions and 17 deletions
+6 -3
View File
@@ -2,10 +2,10 @@ from .state import RbmState
from .matrix import prob, Mat from .matrix import prob, Mat
class EntityParams: class EntityParams:
def __init__(self): def __init__(self, do_gaussian_visible: bool = False, do_gaussian_hidden: bool = False):
# Entity parameters # Entity parameters
self.do_gaussian_visible = False self.do_gaussian_visible = do_gaussian_visible
self.do_gaussian_hidden = False self.do_gaussian_hidden = do_gaussian_hidden
@classmethod @classmethod
def from_dict(cls, params: dict): def from_dict(cls, params: dict):
@@ -24,6 +24,9 @@ class Entity:
self.state = RbmState.from_layer_params(shape) self.state = RbmState.from_layer_params(shape)
self.grad = 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): def grad_zero(self):
self.grad = RbmState.from_layer_params(self.shape) self.grad = RbmState.from_layer_params(self.shape)
+45
View File
@@ -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
+25 -14
View File
@@ -4,18 +4,29 @@ from .entity import Entity
from .status import Status from .status import Status
class TrainingParams: 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 # Training parameters
self.learning_rate = 0.1 self.learning_rate = learning_rate
self.momentum = 0.5 self.momentum = momentum
self.weight_decay = 0 self.weight_decay = weight_decay
self.num_epochs = 1000 self.num_epochs = num_epochs
self.mini_batch_size = 0 self.num_gibbs_samples = num_gibbs_samples
self.do_rao_blackwell = False self.mini_batch_size = mini_batch_size
self.do_gibbs_sample_visible = False self.do_rao_blackwell = do_rao_blackwell
self.do_gibbs_sample_hidden = False self.do_gibbs_sample_visible = do_gibbs_sample_visible
self.do_batch_sample = False self.do_gibbs_sample_hidden = do_gibbs_sample_hidden
self.num_gibbs_samples = 1 self.do_batch_sample = do_batch_sample
@classmethod @classmethod
def from_dict(cls, params: dict): def from_dict(cls, params: dict):
@@ -24,12 +35,12 @@ class TrainingParams:
obj.momentum = params["momentum"] obj.momentum = params["momentum"]
obj.weight_decay = params["weightDecay"] obj.weight_decay = params["weightDecay"]
obj.num_epochs = params["numEpochs"] obj.num_epochs = params["numEpochs"]
obj.num_gibbs_samples = params["numGibbs"]
obj.mini_batch_size = params["miniBatchSize"] obj.mini_batch_size = params["miniBatchSize"]
obj.do_rao_blackwell = params["doRaoBlackwell"] obj.do_rao_blackwell = params["doRaoBlackwell"]
obj.do_gibbs_sample_visible = params["gibbsDoSampleVisible"] obj.do_gibbs_sample_visible = params["gibbsDoSampleVisible"]
obj.do_gibbs_sample_hidden = params["gibbsDoSampleHidden"] obj.do_gibbs_sample_hidden = params["gibbsDoSampleHidden"]
obj.do_batch_sample = params["doSampleBatch"] obj.do_batch_sample = params["doSampleBatch"]
obj.num_gibbs_samples = params["numGibbs"]
return obj return obj
@@ -70,7 +81,7 @@ def cd_jens(entity: Entity, v_states: Mat, params: TrainingParams):
dbv -= np.sum(v_probs, 0) dbv -= np.sum(v_probs, 0)
dbh -= np.sum(h_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): def to_mini_batch(batch: Mat, mini_batch_size: int):
mini_batches = [] mini_batches = []
@@ -97,7 +108,7 @@ def train(entity: Entity, batch: Mat, params: TrainingParams, status: Status, cd
entity.grad_zero() entity.grad_zero()
for epochs in range(params.num_epochs): for epochs in range(params.num_epochs):
# Contrastive divergence learning: calculate gradients # 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 # 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) grad = entity.grad_compute(dbv, dbh, dwhv, learning_rate=params.learning_rate/batch.shape[0], momentum=params.momentum, weight_decay=params.weight_decay)
+40
View File
@@ -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)