- fixed crash if weight not exists
- optimized train loop - choose reasonable default params
This commit is contained in:
+18
-22
@@ -20,25 +20,26 @@ class RbmLayer:
|
||||
self.state.to_file(self.state_filename)
|
||||
|
||||
def load(self):
|
||||
self.state = RbmState.from_file(self.state_filename)
|
||||
state = RbmState.from_file(self.state_filename)
|
||||
if state is not None:
|
||||
self.state = state
|
||||
|
||||
def train(self, batch: np.ndarray, cd_func: Callable, status: Status):
|
||||
training_size = batch.shape[0]
|
||||
num_cases = min(self.params.mini_batch_size, training_size)
|
||||
d_progress = 100.0 / (training_size/num_cases * self.params.num_epochs)
|
||||
progress = 0
|
||||
last_progress = progress
|
||||
training_remain = batch.shape[0]
|
||||
batch_size = min(self.params.mini_batch_size, training_remain)
|
||||
if batch_size == 0:
|
||||
batch_size = training_remain
|
||||
|
||||
d_progress = 100.0 / (training_remain/batch_size * self.params.num_epochs)
|
||||
last_progress = 0
|
||||
batch_row_index = 0
|
||||
|
||||
keep_running = True
|
||||
|
||||
training_remain = training_size
|
||||
training_seen = 0
|
||||
keep_running = True
|
||||
while training_remain > 0 and keep_running:
|
||||
mini_batch_size = min(self.params.mini_batch_size, training_remain)
|
||||
mini_batch = batch[batch_row_index:batch_row_index + mini_batch_size]
|
||||
training_remain -= mini_batch_size
|
||||
batch_row_index += mini_batch_size
|
||||
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)
|
||||
@@ -53,7 +54,7 @@ class RbmLayer:
|
||||
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/num_cases
|
||||
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
|
||||
@@ -62,8 +63,7 @@ class RbmLayer:
|
||||
self.state.b_h += inc_bh
|
||||
self.state.w_hv += inc_whv
|
||||
|
||||
progress = round(epochs*d_progress)
|
||||
|
||||
progress = round(training_seen*d_progress)
|
||||
if progress != last_progress:
|
||||
# Calculate error
|
||||
status.progress = round(progress)
|
||||
@@ -73,7 +73,7 @@ class RbmLayer:
|
||||
break
|
||||
|
||||
last_progress = progress
|
||||
training_seen += 1
|
||||
training_seen += 1
|
||||
|
||||
status.on_change()
|
||||
|
||||
@@ -113,10 +113,6 @@ def xor():
|
||||
params = RbmParams()
|
||||
params.do_rao_blackwell = True
|
||||
params.num_gibbs_samples = 3
|
||||
params.mini_batch_size = 100
|
||||
params.learning_rate = 0.1
|
||||
params.momentum = 0.5
|
||||
params.num_epochs = 1000
|
||||
|
||||
# Create layer
|
||||
layer = RbmLayer("Layer_0", 3, 16, params)
|
||||
|
||||
+4
-4
@@ -11,12 +11,12 @@ class Params:
|
||||
|
||||
class RbmParams(Params):
|
||||
def __init__(self):
|
||||
self.learning_rate = 0
|
||||
self.momentum = 0
|
||||
self.learning_rate = 0.1
|
||||
self.momentum = 0.5
|
||||
self.weight_decay = 0
|
||||
self.num_epochs = 1
|
||||
self.num_epochs = 1000
|
||||
self.num_gibbs_samples = 1
|
||||
self.mini_batch_size = 1
|
||||
self.mini_batch_size = 0
|
||||
|
||||
self.do_gaussian_visible = False
|
||||
self.do_gaussian_hidden = False
|
||||
|
||||
+2
-1
@@ -19,16 +19,17 @@ class RbmState:
|
||||
|
||||
@classmethod
|
||||
def from_file(cls, filename: str):
|
||||
obj = None
|
||||
try:
|
||||
with np.load(filename) as X:
|
||||
w_hv, b_v, b_h = [X[i] for i in ('whv', 'bv', 'bh')]
|
||||
obj = cls(w_hv, b_v, b_h)
|
||||
print(f"{filename} loaded successfully!")
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
except KeyError:
|
||||
pass
|
||||
|
||||
obj = cls(w_hv, b_v, b_h)
|
||||
return obj
|
||||
|
||||
def v_to_h(self, visible: np.ndarray) -> np.ndarray:
|
||||
|
||||
Reference in New Issue
Block a user