- refactored training parameters for grad_compute()
- added l1-norm
- improved tests
This commit is contained in:
2026-01-10 15:21:40 +01:00
parent 00fe5fd178
commit 1722a68b2a
5 changed files with 36 additions and 20 deletions
+22 -10
View File
@@ -7,6 +7,7 @@ class TrainingParams:
learning_rate: float = 0.1,
momentum: float = 0.5,
weight_decay: float = 0.0,
l1_lambda: float = 0.0,
l2_lambda: float = 0.0,
num_epochs: int = 1000,
num_gibbs_samples: int = 1,
@@ -20,6 +21,7 @@ class TrainingParams:
self.learning_rate = learning_rate
self.momentum = momentum
self.weight_decay = weight_decay
self.l1_lambda = l1_lambda
self.l2_lambda = l2_lambda
self.num_epochs = num_epochs
self.num_gibbs_samples = num_gibbs_samples
@@ -86,22 +88,32 @@ class Entity:
def grad_zero(self):
self.grad = RbmState.from_layer_params(self.shape)
def grad_compute(self, d_bv: Mat, d_bh: Mat, d_whv: Mat, learning_rate: float, momentum: float, weight_decay: float, l2_lambda: float):
def grad_compute(self, d_bv: Mat, d_bh: Mat, d_whv: Mat):
params = self.training_params
lr = params.learning_rate
# 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.b_v = params.momentum * self.grad.b_v + lr * d_bv
self.grad.b_h = params.momentum * self.grad.b_h + lr * d_bh
# compute L1 term and penalize cost function (d_whv)
l1_norm = np.sum(np.abs(self.state.w_hv))
l1_term = params.l1_lambda*l1_norm*self.state.w_hv
# compute L2 term and penalize cost function (d_whv)
l2_norm = 0.5*np.sum(np.square(self.state.w_hv))
l2_term = l2_lambda*l2_norm*self.state.w_hv
self.grad.w_hv = momentum * self.grad.w_hv + learning_rate * (d_whv-l2_term) - learning_rate * weight_decay * self.state.w_hv
l2_norm = np.sum(np.square(self.state.w_hv))
l2_term = params.l2_lambda*l2_norm*self.state.w_hv
self.grad.w_hv = params.momentum * self.grad.w_hv + lr * (d_whv-(l1_term+l2_term)) - lr * params.weight_decay * self.state.w_hv
return self.grad
def state_adjust(self, grad: RbmState):
def state_adjust(self, grad: RbmState, k: float = 1.0):
# Adjust state
self.state.b_v += grad.b_v
self.state.b_h += grad.b_h
self.state.w_hv += grad.w_hv
self.state.b_v += k*grad.b_v
self.state.b_h += k*grad.b_h
self.state.w_hv += k*grad.w_hv
def forward(self, v: Mat, num_gibbs: int = 0) -> Mat:
num_gibbs = self.params.num_gibbs_samples if num_gibbs == 0 else num_gibbs
+1 -2
View File
@@ -35,8 +35,7 @@ class Optimizer:
dwhv, dbv, dbh = self.loss(self.entity, data)
# Adjust weight and biases
grad = self.entity.grad_compute(dbv, dbh, dwhv, learning_rate=params.learning_rate / data.shape[0],
momentum=params.momentum, weight_decay=params.weight_decay)
grad = self.entity.grad_compute(dbv, dbh, dwhv)
# Adjust weights
self.entity.state_adjust(grad)
+2 -2
View File
@@ -225,8 +225,8 @@ def train(entity: Entity, batch: Mat, status: Status):
dwhv, dbv, dbh = cd_func(entity, mini_batch)
# Adjust weight and biases
grad = entity.grad_compute(dbv, dbh, dwhv, learning_rate=params.learning_rate/batch.shape[0], momentum=params.momentum, weight_decay=params.weight_decay, l2_lambda=params.l2_lambda)
entity.state_adjust(grad)
grad = entity.grad_compute(dbv, dbh, dwhv)
entity.state_adjust(grad, 1.0/mini_batch.shape[0])
# check if status update is needed
if status.want_report(round(progress)):