refactored
This commit is contained in:
+15
-12
@@ -22,21 +22,24 @@ class Entity:
|
|||||||
self.shape = shape
|
self.shape = shape
|
||||||
self.params = params
|
self.params = params
|
||||||
self.state = RbmState.from_layer_params(shape)
|
self.state = RbmState.from_layer_params(shape)
|
||||||
self.delta_state = RbmState.from_layer_params(shape)
|
self.grad = RbmState.from_layer_params(shape)
|
||||||
|
|
||||||
def prepare(self):
|
def grad_zero(self):
|
||||||
self.delta_state = RbmState.from_layer_params(self.shape)
|
self.grad = RbmState.from_layer_params(self.shape)
|
||||||
|
|
||||||
def adjust(self, d_bv: Mat, d_bh: Mat, d_whv: Mat, learning_rate: float, momentum: float, weight_decay: float):
|
def grad_compute(self, d_bv: Mat, d_bh: Mat, d_whv: Mat, learning_rate: float, momentum: float, weight_decay: float):
|
||||||
# Create delta
|
# Compute gradient
|
||||||
self.delta_state.b_v = (momentum * self.delta_state.b_v + learning_rate*d_bv)
|
self.grad.b_v = (momentum * self.grad.b_v + learning_rate * d_bv)
|
||||||
self.delta_state.b_h = (momentum * self.delta_state.b_h + learning_rate*d_bh)
|
self.grad.b_h = (momentum * self.grad.b_h + learning_rate * d_bh)
|
||||||
self.delta_state.w_hv = (momentum * self.delta_state.w_hv + learning_rate*d_whv - weight_decay*self.state.w_hv)
|
self.grad.w_hv = (momentum * self.grad.w_hv + learning_rate * d_whv - weight_decay * self.state.w_hv)
|
||||||
|
|
||||||
# Create Adjust
|
return self.grad
|
||||||
self.state.b_v += self.delta_state.b_v
|
|
||||||
self.state.b_h += self.delta_state.b_h
|
def state_adjust(self, grad: RbmState):
|
||||||
self.state.w_hv += self.delta_state.w_hv
|
# Adjust state
|
||||||
|
self.state.b_v += grad.b_v
|
||||||
|
self.state.b_h += grad.b_h
|
||||||
|
self.state.w_hv += grad.w_hv
|
||||||
|
|
||||||
def forward(self, v: Mat, num_gibbs: int = 1) -> Mat:
|
def forward(self, v: Mat, num_gibbs: int = 1) -> Mat:
|
||||||
h = self._v_to_ph(v)
|
h = self._v_to_ph(v)
|
||||||
|
|||||||
+3
-2
@@ -94,13 +94,14 @@ def train(entity: Entity, batch: Mat, params: TrainingParams, status: Status, cd
|
|||||||
if not keep_running:
|
if not keep_running:
|
||||||
break
|
break
|
||||||
|
|
||||||
entity.prepare()
|
entity.grad_zero()
|
||||||
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, mini_batch, params)
|
dwhv, dbv, dbh, _ = cd_func(entity, mini_batch, params)
|
||||||
|
|
||||||
# Adjust weight and biases
|
# Adjust weight and biases
|
||||||
entity.adjust(dbv, dbh, dwhv, learning_rate=params.learning_rate/batch.shape[0], momentum=params.momentum, weight_decay=params.weight_decay)
|
grad = entity.grad_compute(dbv, dbh, dwhv, learning_rate=params.learning_rate/batch.shape[0], momentum=params.momentum, weight_decay=params.weight_decay)
|
||||||
|
entity.state_adjust(grad)
|
||||||
|
|
||||||
# check if status update is needed
|
# check if status update is needed
|
||||||
if status.want_report(round(progress)):
|
if status.want_report(round(progress)):
|
||||||
|
|||||||
Reference in New Issue
Block a user