- numGibbs no longer part of EntityParameter

- refactored forward and reconstruct
- conditionally use optimizer for training
This commit is contained in:
2025-12-19 17:15:10 +01:00
parent a47922cb1c
commit d3c9fe4681
4 changed files with 124 additions and 41 deletions
+18 -21
View File
@@ -4,15 +4,12 @@ from .matrix import prob, Mat
class EntityParams: class EntityParams:
def __init__(self): def __init__(self):
# Entity parameters # Entity parameters
self.num_gibbs_samples = 1
self.do_gaussian_visible = False self.do_gaussian_visible = False
self.do_gaussian_hidden = False self.do_gaussian_hidden = False
@classmethod @classmethod
def from_dict(cls, params: dict): def from_dict(cls, params: dict):
obj = EntityParams() obj = EntityParams()
if "numGibbs" in params:
obj.num_gibbs_samples = params["numGibbs"]
if "doGaussianVisible" in params: if "doGaussianVisible" in params:
obj.do_gaussian_visible = params["doGaussianVisible"] obj.do_gaussian_visible = params["doGaussianVisible"]
if "doGaussianHidden" in params: if "doGaussianHidden" in params:
@@ -26,35 +23,35 @@ class Entity:
self.params = params self.params = params
self.state = RbmState.from_layer_params(shape) self.state = RbmState.from_layer_params(shape)
def v_to_ph(self, v: Mat) -> Mat: def forward(self, v: Mat, num_gibbs: int = 1) -> Mat:
h = self._v_to_ph(v)
for i in range(num_gibbs-1):
h = self._h_to_pv(h)
h = self._v_to_ph(h)
return h
def reconstruct(self, h: Mat, num_gibbs: int = 1) -> Mat:
v = self._h_to_pv(h)
for i in range(num_gibbs-1):
v = self._v_to_ph(v)
v = self._h_to_pv(v)
return v
def _v_to_ph(self, v: Mat) -> Mat:
state = self.state.v_to_h(v) state = self.state.v_to_h(v)
if self.params.do_gaussian_hidden: if self.params.do_gaussian_hidden:
return state return state
return prob(state) return prob(state)
def h_to_pv(self, h: Mat) -> Mat: def _h_to_pv(self, h: Mat) -> Mat:
state = self.state.h_to_v(h) state = self.state.h_to_v(h)
if self.params.do_gaussian_visible: if self.params.do_gaussian_visible:
return state return state
return prob(state) return prob(state)
def forward(self, v: Mat) -> Mat:
h = self.v_to_ph(v)
for i in range(self.params.num_gibbs_samples-1):
h = self.h_to_pv(h)
h = self.v_to_ph(h)
return h
def reconstruct(self, h: Mat) -> Mat:
v = self.h_to_pv(h)
for i in range(self.params.num_gibbs_samples-1):
v = self.v_to_ph(v)
v = self.h_to_pv(v)
return v
if __name__ == "__main__": if __name__ == "__main__":
print("Test: [passed]") print("Test: [passed]")
+10 -4
View File
@@ -1,8 +1,10 @@
from .status import Status from .status import Status
from .train import train from .train import train, Optimizer
from .stack import Stack, StackType from .stack import Stack, StackType
from .matrix import Mat, np from .matrix import Mat, np
USE_OPTIMIZER = False
class StackDeep(Stack): class StackDeep(Stack):
def __init__(self, name: str, work_dir: str = '.'): def __init__(self, name: str, work_dir: str = '.'):
Stack.__init__(self, StackType.Deep, name, work_dir) Stack.__init__(self, StackType.Deep, name, work_dir)
@@ -12,7 +14,7 @@ class StackDeep(Stack):
for index, layer in enumerate(self.layers): for index, layer in enumerate(self.layers):
if index == from_layer_id: if index == from_layer_id:
break break
_batch = layer.entity.v_to_ph(_batch) _batch = layer.entity.forward(_batch)
return _batch return _batch
@@ -20,8 +22,12 @@ class StackDeep(Stack):
_batch = np.copy(batch) _batch = np.copy(batch)
for index, layer in enumerate(self.layers): for index, layer in enumerate(self.layers):
print(f"Train layer {index} for {layer.training_params.num_epochs} epochs") print(f"Train layer {index} for {layer.training_params.num_epochs} epochs")
_batch = self.batch_from(batch, index) if USE_OPTIMIZER:
train(layer.entity, _batch, layer.training_params, status=status) optim = Optimizer(layer.entity, layer.training_params)
_batch = optim(_batch, status=status)
else:
_batch = self.batch_from(batch, index)
train(layer.entity, _batch, layer.training_params, status=status)
def pass_up(self, visible: Mat, from_layer_id: int = 0): def pass_up(self, visible: Mat, from_layer_id: int = 0):
+85 -11
View File
@@ -15,6 +15,7 @@ class TrainingParams:
self.do_gibbs_sample_visible = False self.do_gibbs_sample_visible = False
self.do_gibbs_sample_hidden = False self.do_gibbs_sample_hidden = False
self.do_batch_sample = False self.do_batch_sample = False
self.num_gibbs_samples = 1
@classmethod @classmethod
def from_dict(cls, params: dict): def from_dict(cls, params: dict):
@@ -28,18 +29,19 @@ class TrainingParams:
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
def cd_jens(entity: Entity, v_states: Mat, params: TrainingParams): def cd_jens(entity: Entity, v_states: Mat, params: TrainingParams):
v_probs = prob(v_states) v_probs = prob(v_states)
h_states = entity.v_to_ph(v_states) h_states = entity.forward(v_states)
h_probs = h_states h_probs = h_states
if entity.params.do_gaussian_hidden: if entity.params.do_gaussian_hidden:
h_states += gaussian(h_states.shape) h_states += gaussian(h_states.shape)
else: else:
h_probs = entity.v_to_ph(v_states) h_probs = entity.forward(v_states)
if params.do_rao_blackwell: if params.do_rao_blackwell:
h_states = h_probs h_states = h_probs
else: else:
@@ -51,24 +53,24 @@ def cd_jens(entity: Entity, v_states: Mat, params: TrainingParams):
dbh = np.sum(h_states, 0) dbh = np.sum(h_states, 0)
# Gibbs sampling with training params # Gibbs sampling with training params
for i in range(entity.params.num_gibbs_samples): for i in range(params.num_gibbs_samples):
if params.do_gibbs_sample_hidden: if params.do_gibbs_sample_hidden:
v_probs = entity.h_to_pv(sample(h_probs)) v_probs = entity.reconstruct(sample(h_probs))
else: else:
v_probs = entity.h_to_pv(h_probs) v_probs = entity.reconstruct(h_probs)
# Create hidden representation given v # Create hidden representation given v
if params.do_gibbs_sample_visible: if params.do_gibbs_sample_visible:
h_probs = entity.v_to_ph(sample(v_probs)) h_probs = entity.forward(sample(v_probs))
else: else:
h_probs = entity.v_to_ph(v_probs) h_probs = entity.forward(v_probs)
# Update weights (negative phase) # Update weights (negative phase)
dw -= np.dot(np.transpose(v_probs), h_probs) dw -= np.dot(np.transpose(v_probs), h_probs)
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 return dw, dbv, dbh, h_probs
def train(entity: Entity, batch: Mat, params: TrainingParams, status: Status, cd_func: Callable = cd_jens): def train(entity: Entity, batch: Mat, params: TrainingParams, status: Status, cd_func: Callable = cd_jens):
training_remain = batch.shape[0] training_remain = batch.shape[0]
@@ -97,7 +99,7 @@ def train(entity: Entity, batch: Mat, params: TrainingParams, status: Status, cd
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, v_states, params) dwhv, dbv, dbh, _ = cd_func(entity, v_states, params)
# Adjust weight and biases # Adjust weight and biases
kl = params.learning_rate/batch_size kl = params.learning_rate/batch_size
@@ -112,7 +114,7 @@ def train(entity: Entity, batch: Mat, params: TrainingParams, status: Status, cd
# check if status update is needed # check if status update is needed
if status.want_report(round(progress)): if status.want_report(round(progress)):
# Calculate error # Calculate error
err_rms = rms_error_accu(mini_batch - entity.h_to_pv(entity.v_to_ph(v_states))) err_rms = rms_error_accu(batch - entity.reconstruct(entity.forward(batch)))
if not status.on_change({"progress": {"value": round(progress), "unit": "%"}, if not status.on_change({"progress": {"value": round(progress), "unit": "%"},
"err_rms": {"value": err_rms, "unit": ""}}): "err_rms": {"value": err_rms, "unit": ""}}):
keep_running = False keep_running = False
@@ -121,6 +123,78 @@ def train(entity: Entity, batch: Mat, params: TrainingParams, status: Status, cd
progress += d_progress*batch_size_remain progress += d_progress*batch_size_remain
# Update final status # Update final status
err_rms = rms_error_accu(batch - entity.h_to_pv(entity.v_to_ph(batch))) err_rms = rms_error_accu(batch - entity.reconstruct(entity.forward(batch)))
status.on_change({"progress": {"value": round(progress), "unit": "%"}, status.on_change({"progress": {"value": round(progress), "unit": "%"},
"err_rms_total": {"value": err_rms, "unit": ""}}) "err_rms_total": {"value": err_rms, "unit": ""}})
class Optimizer:
def __init__(self, entity: Entity, params: TrainingParams, cd_func: Callable = cd_jens):
self.inc_bv = np.zeros(entity.state.b_v.shape)
self.inc_bh = np.zeros(entity.state.b_h.shape)
self.inc_whv = np.zeros(entity.state.w_hv.shape)
self.entity = entity
self.params = params
self.cd = cd_func
def __call__(self, batch: Mat, status: Status):
return self.process(batch, status)
def process(self, batch: Mat, status: Status):
status.on_change()
d_progress = 100.0 / self.params.num_epochs
progress = 0
v_states = batch
if self.params.do_batch_sample:
v_states = sample(batch)
forward = None
for epochs in range(self.params.num_epochs):
forward = self.epoch(v_states)
# check if status update is needed
if status.want_report(round(progress)):
# Calculate error
result = self.entity.reconstruct(forward)
err_rms = rms_error_accu(batch - result)
if not status.on_change({"progress": {"value": round(progress), "unit": "%"},
"err_rms": {"value": err_rms, "unit": ""}}):
break
progress += d_progress
err_rms = rms_error_accu(batch - self.entity.reconstruct(forward))
status.on_change({"progress": {"value": round(progress), "unit": "%"},
"err_rms": {"value": err_rms, "unit": ""}})
return forward
def epoch(self, batch: Mat):
forward = None
training_remain = batch.shape[0]
batch_row_index = 0
while training_remain > 0:
batch_size = min(self.params.mini_batch_size, training_remain)
if batch_size == 0:
batch_size = training_remain
mini_batch = batch[batch_row_index:batch_row_index + batch_size]
# Contrastive divergence learning: calculate gradients
dwhv, dbv, dbh, forward = self.cd(self.entity, mini_batch, self.params)
# Adjust weight and biases
kl = self.params.learning_rate / batch_size
inc_bv = self.params.momentum * self.inc_bv + kl * dbv
inc_bh = self.params.momentum * self.inc_bh + kl * dbh
inc_whv = self.params.momentum * self.inc_whv + kl * dwhv - self.params.weight_decay * self.entity.state.w_hv
self.entity.state.b_v += inc_bv
self.entity.state.b_h += inc_bh
self.entity.state.w_hv += inc_whv
training_remain -= batch_size
batch_row_index += batch_size
return forward
+11 -5
View File
@@ -2,11 +2,13 @@ import os.path
from rbm.layer import Layer from rbm.layer import Layer
from rbm.status import Status from rbm.status import Status
from rbm.train import train, TrainingParams from rbm.train import train, TrainingParams, Optimizer
from rbm.matrix import Mat, np from rbm.matrix import Mat, np
from rbm.entity import EntityParams from rbm.entity import EntityParams
work_dir = "../../results" WORK_DIR = "../../results"
USE_OPTIMIZER = True
def xor(): def xor():
# Create params # Create params
entity_params = EntityParams() entity_params = EntityParams()
@@ -21,16 +23,20 @@ def xor():
layer.init(0.01) layer.init(0.01)
# Load weights (if exists) # Load weights (if exists)
layer.load(os.path.join(work_dir, "xor_layer0_state.npz")) layer.load(os.path.join(WORK_DIR, "xor_layer0_state.npz"))
# Prepare training data # Prepare training data
training_batch = Mat([[0,1,1], [0,0,0], [1,1,0], [1,0,1]], dtype=np.float64) training_batch = Mat([[0,1,1], [0,0,0], [1,1,0], [1,0,1]], dtype=np.float64)
# Train layer # Train layer
train(layer.entity, training_batch, training_params, Status()) if USE_OPTIMIZER:
optim = Optimizer(layer.entity, training_params)
optim(training_batch, Status())
else:
train(layer.entity, training_batch, training_params, Status())
# Save weights # Save weights
layer.save(os.path.join(work_dir, "xor_layer0_state.npz")) layer.save(os.path.join(WORK_DIR, "xor_layer0_state.npz"))
# Test with test data # Test with test data
test_batch = Mat([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64) test_batch = Mat([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64)