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
+25 -14
View File
@@ -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)