refactored

This commit is contained in:
2025-12-21 10:33:04 +01:00
parent 8ed6488b45
commit 7668c73ea4
2 changed files with 18 additions and 14 deletions
+15 -12
View File
@@ -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)