refactored train
This commit is contained in:
@@ -22,6 +22,21 @@ class Entity:
|
||||
self.shape = shape
|
||||
self.params = params
|
||||
self.state = RbmState.from_layer_params(shape)
|
||||
self.delta_state = RbmState.from_layer_params(shape)
|
||||
|
||||
def prepare(self):
|
||||
self.delta_state = 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)
|
||||
|
||||
# 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
|
||||
|
||||
def forward(self, v: Mat, num_gibbs: int = 1) -> Mat:
|
||||
h = self._v_to_ph(v)
|
||||
|
||||
Reference in New Issue
Block a user