refactored

This commit is contained in:
2025-12-18 12:24:17 +01:00
parent 470d2bfa9a
commit 5cbef6b09f
3 changed files with 102 additions and 90 deletions
+94
View File
@@ -0,0 +1,94 @@
import numpy as np
from collections.abc import Callable
from params import RbmParams
from state import RbmState
from matrix import sample, prob, rms_error_accu
from status import Status
class Entity:
def __init__(self, shape: tuple[int, int], params: RbmParams):
self.state = RbmState.from_layer_params(shape)
self.params = params
def train(self, batch: np.ndarray, cd_func: Callable, status: Status):
training_remain = batch.shape[0]
batch_size = min(self.params.mini_batch_size, training_remain)
if batch_size == 0:
batch_size = training_remain
status.on_change({})
d_progress = 100.0 / (training_remain/batch_size * self.params.num_epochs)
batch_row_index = 0
training_seen = 0
keep_running = True
while training_remain > 0 and keep_running:
batch_size_remain = min(batch_size, training_remain)
mini_batch = batch[batch_row_index:batch_row_index + batch_size_remain]
training_remain -= batch_size_remain
batch_row_index += batch_size_remain
inc_bv = np.zeros(self.state.b_v.shape)
inc_bh = np.zeros(self.state.b_h.shape)
inc_whv = np.zeros(self.state.w_hv.shape)
v_states = mini_batch
if self.params.do_batch_sample:
v_states = sample(mini_batch)
for epochs in range(self.params.num_epochs):
# Contrastive divergence learning: calculate gradients
dwhv, dbv, dbh = cd_func(v_states, self.params, self.v_to_ph, self.h_to_pv)
# Adjust weight and biases
kl = self.params.learning_rate/batch_size
inc_bv = self.params.momentum*inc_bv + kl*dbv
inc_bh = self.params.momentum*inc_bh + kl*dbh
inc_whv = self.params.momentum*inc_whv + kl*dwhv - self.params.weight_decay*self.state.w_hv
self.state.b_v += inc_bv
self.state.b_h += inc_bh
self.state.w_hv += inc_whv
# Calculate error
if status.want_report(round(training_seen*d_progress)):
err_rms = rms_error_accu(mini_batch - self.h_to_pv(self.v_to_ph(v_states)))
if not status.on_change({"progress": {"value": round(training_seen * d_progress), "unit": "%"},
"err_rms": {"value": err_rms, "unit": ""}}):
keep_running = False
break
training_seen += 1
def v_to_ph(self, v: np.ndarray) -> np.ndarray:
state = self.state.v_to_h(v)
if self.params.do_gaussian_visible:
return state
return prob(state)
def h_to_pv(self, h: np.ndarray) -> np.ndarray:
state = self.state.h_to_v(h)
if self.params.do_gaussian_visible:
return state
return prob(state)
def gibbs_v_to_h(self, v: np.ndarray) -> np.ndarray:
h = None
for i in range(self.params.num_gibbs_samples):
h = self.v_to_ph(v)
v = self.h_to_pv(h)
return h
def gibbs_h_to_v(self, h: np.ndarray) -> np.ndarray:
v = None
for i in range(self.params.num_gibbs_samples):
v = self.h_to_pv(h)
h = self.v_to_ph(v)
return v
if __name__ == "__main__":
print("Test: [passed]")
+8 -90
View File
@@ -1,113 +1,31 @@
import numpy as np import numpy as np
from collections.abc import Callable
from params import RbmParams from params import RbmParams
from state import RbmState from state import RbmState
from matrix import sample, prob, rms_error_accu
from status import Status from status import Status
from cd_train import cd_jens from cd_train import cd_jens
from entity import Entity
class Layer: class Layer:
def __init__(self, name: str, shape: tuple[int, int, int, int], params: RbmParams): def __init__(self, name: str, shape: tuple[int, int, int, int], params: RbmParams):
self.name = name self.name = name
self.shape = shape self.shape = shape
self.state = RbmState.from_layer_params((shape[0]*shape[1]+shape[2], shape[3])) self.entity = Entity((shape[0]*shape[1]+shape[2], shape[3]), params)
self.params = params
self.state_filename = f"{self.name}_state.npz" self.state_filename = f"{self.name}_state.npz"
def init(self, std: float): def init(self, std: float):
self.state.init(mu=0, std=std) self.entity.state.init(mu=0, std=std)
def save(self, filename: str = None): def save(self, filename: str = None):
if filename is None: if filename is None:
filename = self.state_filename filename = self.state_filename
self.state.to_file(filename) self.entity.state.to_file(filename)
def load(self, filename: str = None): def load(self, filename: str = None):
if filename is None: if filename is None:
filename = self.state_filename filename = self.state_filename
state = RbmState.from_file(filename) state = RbmState.from_file(filename)
if state is not None: if state is not None:
self.state = state self.entity.state = state
def train(self, batch: np.ndarray, cd_func: Callable, status: Status):
training_remain = batch.shape[0]
batch_size = min(self.params.mini_batch_size, training_remain)
if batch_size == 0:
batch_size = training_remain
status.on_change({})
d_progress = 100.0 / (training_remain/batch_size * self.params.num_epochs)
batch_row_index = 0
training_seen = 0
keep_running = True
while training_remain > 0 and keep_running:
batch_size_remain = min(batch_size, training_remain)
mini_batch = batch[batch_row_index:batch_row_index + batch_size_remain]
training_remain -= batch_size_remain
batch_row_index += batch_size_remain
inc_bv = np.zeros(self.state.b_v.shape)
inc_bh = np.zeros(self.state.b_h.shape)
inc_whv = np.zeros(self.state.w_hv.shape)
v_states = mini_batch
if self.params.do_batch_sample:
v_states = sample(mini_batch)
for epochs in range(self.params.num_epochs):
# Contrastive divergence learning: calculate gradients
dwhv, dbv, dbh = cd_func(v_states, self.params, self.v_to_ph, self.h_to_pv)
# Adjust weight and biases
kl = self.params.learning_rate/batch_size
inc_bv = self.params.momentum*inc_bv + kl*dbv
inc_bh = self.params.momentum*inc_bh + kl*dbh
inc_whv = self.params.momentum*inc_whv + kl*dwhv - self.params.weight_decay*self.state.w_hv
self.state.b_v += inc_bv
self.state.b_h += inc_bh
self.state.w_hv += inc_whv
# Calculate error
if status.want_report(round(training_seen*d_progress)):
err_rms = rms_error_accu(mini_batch - self.h_to_pv(self.v_to_ph(v_states)))
if not status.on_change({"progress": {"value": round(training_seen * d_progress), "unit": "%"},
"err_rms": {"value": err_rms, "unit": ""}}):
keep_running = False
break
training_seen += 1
def v_to_ph(self, v: np.ndarray) -> np.ndarray:
state = self.state.v_to_h(v)
if self.params.do_gaussian_visible:
return state
return prob(state)
def h_to_pv(self, h: np.ndarray) -> np.ndarray:
state = self.state.h_to_v(h)
if self.params.do_gaussian_visible:
return state
return prob(state)
def gibbs_v_to_h(self, v: np.ndarray) -> np.ndarray:
h = None
for i in range(self.params.num_gibbs_samples):
h = self.v_to_ph(v)
v = self.h_to_pv(h)
return h
def gibbs_h_to_v(self, h: np.ndarray) -> np.ndarray:
v = None
for i in range(self.params.num_gibbs_samples):
v = self.h_to_pv(h)
h = self.v_to_ph(v)
return v
def xor(): def xor():
# Create params # Create params
@@ -128,7 +46,7 @@ def xor():
training_batch = np.array([[0,1,1], [0,0,0], [1,1,0], [1,0,1]], dtype=np.float64) training_batch = np.array([[0,1,1], [0,0,0], [1,1,0], [1,0,1]], dtype=np.float64)
# Train layer # Train layer
layer.train(training_batch, cd_jens, Status()) layer.entity.train(training_batch, cd_jens, Status())
# Save weights # Save weights
layer.save() layer.save()
@@ -136,8 +54,8 @@ def xor():
# Test with test data # Test with test data
test_batch = np.array([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64) test_batch = np.array([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64)
for pattern in test_batch: for pattern in test_batch:
h = layer.gibbs_v_to_h(pattern) h = layer.entity.gibbs_v_to_h(pattern)
v = layer.gibbs_h_to_v(h) v = layer.entity.gibbs_h_to_v(h)
print(f"P{pattern} : {v}") print(f"P{pattern} : {v}")