- training params are (again) attrubute of Entity
- ditched layout concept
This commit is contained in:
+13
-49
@@ -1,50 +1,9 @@
|
||||
from collections.abc import Callable
|
||||
from .matrix import sample, sample_gaussian, prob, rms_error_accu, Mat, np
|
||||
from .entity import Entity
|
||||
from .status import Status
|
||||
|
||||
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
|
||||
|
||||
def cd_jens(entity: Entity, v_states: Mat, params: TrainingParams):
|
||||
def cd_jens(entity: Entity, v_states: Mat):
|
||||
params = entity.training_params
|
||||
v_probs = prob(v_states)
|
||||
h_states = entity.forward(v_states)
|
||||
h_probs = h_states
|
||||
@@ -83,7 +42,8 @@ def cd_jens(entity: Entity, v_states: Mat, params: TrainingParams):
|
||||
|
||||
return dw, dbv, dbh
|
||||
|
||||
def cd_binary_binary(entity: Entity, data_pos: Mat, params: TrainingParams):
|
||||
def cd_binary_binary(entity: Entity, data_pos: Mat):
|
||||
params = entity.training_params
|
||||
# Positive phase
|
||||
h_probs_pos = prob(entity.h_given_v(data_pos))
|
||||
|
||||
@@ -112,7 +72,7 @@ def cd_binary_binary(entity: Entity, data_pos: Mat, params: TrainingParams):
|
||||
|
||||
return dw, dbv, dbh
|
||||
|
||||
def cd_gaussian_binary(entity: Entity, data_pos: Mat, params: TrainingParams):
|
||||
def cd_gaussian_binary(entity: Entity, data_pos: Mat):
|
||||
# Positive phase
|
||||
h_probs_pos = entity.h_given_v(data_pos)
|
||||
|
||||
@@ -135,7 +95,7 @@ def cd_gaussian_binary(entity: Entity, data_pos: Mat, params: TrainingParams):
|
||||
|
||||
return dw, dbv, dbh
|
||||
|
||||
def cd_gaussian_gaussian(entity: Entity, data_pos: Mat, params: TrainingParams):
|
||||
def cd_gaussian_gaussian(entity: Entity, data_pos: Mat):
|
||||
# Positive phase
|
||||
h_probs_pos = entity.h_given_v(data_pos)
|
||||
|
||||
@@ -170,7 +130,11 @@ def to_mini_batch(batch: Mat, mini_batch_size: int):
|
||||
|
||||
return mini_batches
|
||||
|
||||
def train(entity: Entity, batch: Mat, params: TrainingParams, status: Status):
|
||||
def train(entity: Entity, batch: Mat, status: Status):
|
||||
params = entity.training_params
|
||||
if params is None:
|
||||
return False
|
||||
|
||||
mini_batch_size = min(params.mini_batch_size, batch.shape[0]) if params.mini_batch_size > 0 else batch.shape[0]
|
||||
d_progress = 100.0 / (batch.shape[0]*params.num_epochs)
|
||||
progress = 0
|
||||
@@ -191,7 +155,7 @@ def train(entity: Entity, batch: Mat, params: TrainingParams, status: Status):
|
||||
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)
|
||||
|
||||
# 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)
|
||||
@@ -213,4 +177,4 @@ def train(entity: Entity, batch: Mat, params: TrainingParams, status: Status):
|
||||
status.on_change({"progress": {"value": round(progress), "unit": "%"},
|
||||
"err_rms_total": {"value": err_rms, "unit": ""}})
|
||||
|
||||
return None
|
||||
return True
|
||||
|
||||
Reference in New Issue
Block a user