diff --git a/src/rbm/entity.py b/src/rbm/entity.py index 78da2a6..df53acf 100644 --- a/src/rbm/entity.py +++ b/src/rbm/entity.py @@ -22,21 +22,24 @@ class Entity: self.shape = shape self.params = params 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): - self.delta_state = RbmState.from_layer_params(self.shape) + def grad_zero(self): + 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): - # Create delta - self.delta_state.b_v = (momentum * self.delta_state.b_v + learning_rate*d_bv) - self.delta_state.b_h = (momentum * self.delta_state.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) + def grad_compute(self, d_bv: Mat, d_bh: Mat, d_whv: Mat, learning_rate: float, momentum: float, weight_decay: float): + # Compute gradient + self.grad.b_v = (momentum * self.grad.b_v + learning_rate * d_bv) + self.grad.b_h = (momentum * self.grad.b_h + learning_rate * d_bh) + self.grad.w_hv = (momentum * self.grad.w_hv + learning_rate * d_whv - weight_decay * self.state.w_hv) - # Create Adjust - self.state.b_v += self.delta_state.b_v - self.state.b_h += self.delta_state.b_h - self.state.w_hv += self.delta_state.w_hv + return self.grad + + def state_adjust(self, grad: RbmState): + # 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: h = self._v_to_ph(v) diff --git a/src/rbm/train.py b/src/rbm/train.py index 83a8a90..23e3942 100644 --- a/src/rbm/train.py +++ b/src/rbm/train.py @@ -94,13 +94,14 @@ def train(entity: Entity, batch: Mat, params: TrainingParams, status: Status, cd if not keep_running: break - entity.prepare() + entity.grad_zero() for epochs in range(params.num_epochs): # Contrastive divergence learning: calculate gradients dwhv, dbv, dbh, _ = cd_func(entity, mini_batch, params) # 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 if status.want_report(round(progress)):