[entity]
- refactored training parameters for grad_compute() - added l1-norm - improved tests
This commit is contained in:
+22
-10
@@ -7,6 +7,7 @@ class TrainingParams:
|
|||||||
learning_rate: float = 0.1,
|
learning_rate: float = 0.1,
|
||||||
momentum: float = 0.5,
|
momentum: float = 0.5,
|
||||||
weight_decay: float = 0.0,
|
weight_decay: float = 0.0,
|
||||||
|
l1_lambda: float = 0.0,
|
||||||
l2_lambda: float = 0.0,
|
l2_lambda: float = 0.0,
|
||||||
num_epochs: int = 1000,
|
num_epochs: int = 1000,
|
||||||
num_gibbs_samples: int = 1,
|
num_gibbs_samples: int = 1,
|
||||||
@@ -20,6 +21,7 @@ class TrainingParams:
|
|||||||
self.learning_rate = learning_rate
|
self.learning_rate = learning_rate
|
||||||
self.momentum = momentum
|
self.momentum = momentum
|
||||||
self.weight_decay = weight_decay
|
self.weight_decay = weight_decay
|
||||||
|
self.l1_lambda = l1_lambda
|
||||||
self.l2_lambda = l2_lambda
|
self.l2_lambda = l2_lambda
|
||||||
self.num_epochs = num_epochs
|
self.num_epochs = num_epochs
|
||||||
self.num_gibbs_samples = num_gibbs_samples
|
self.num_gibbs_samples = num_gibbs_samples
|
||||||
@@ -86,22 +88,32 @@ class Entity:
|
|||||||
def grad_zero(self):
|
def grad_zero(self):
|
||||||
self.grad = RbmState.from_layer_params(self.shape)
|
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
|
# Compute gradient
|
||||||
self.grad.b_v = momentum * self.grad.b_v + learning_rate * d_bv
|
self.grad.b_v = params.momentum * self.grad.b_v + lr * d_bv
|
||||||
self.grad.b_h = momentum * self.grad.b_h + learning_rate * d_bh
|
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)
|
# compute L2 term and penalize cost function (d_whv)
|
||||||
l2_norm = 0.5*np.sum(np.square(self.state.w_hv))
|
l2_norm = np.sum(np.square(self.state.w_hv))
|
||||||
l2_term = l2_lambda*l2_norm*self.state.w_hv
|
l2_term = params.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
|
|
||||||
|
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
|
return self.grad
|
||||||
|
|
||||||
def state_adjust(self, grad: RbmState):
|
def state_adjust(self, grad: RbmState, k: float = 1.0):
|
||||||
# Adjust state
|
# Adjust state
|
||||||
self.state.b_v += grad.b_v
|
self.state.b_v += k*grad.b_v
|
||||||
self.state.b_h += grad.b_h
|
self.state.b_h += k*grad.b_h
|
||||||
self.state.w_hv += grad.w_hv
|
self.state.w_hv += k*grad.w_hv
|
||||||
|
|
||||||
def forward(self, v: Mat, num_gibbs: int = 0) -> Mat:
|
def forward(self, v: Mat, num_gibbs: int = 0) -> Mat:
|
||||||
num_gibbs = self.params.num_gibbs_samples if num_gibbs == 0 else num_gibbs
|
num_gibbs = self.params.num_gibbs_samples if num_gibbs == 0 else num_gibbs
|
||||||
|
|||||||
+1
-2
@@ -35,8 +35,7 @@ class Optimizer:
|
|||||||
dwhv, dbv, dbh = self.loss(self.entity, data)
|
dwhv, dbv, dbh = self.loss(self.entity, data)
|
||||||
|
|
||||||
# Adjust weight and biases
|
# Adjust weight and biases
|
||||||
grad = self.entity.grad_compute(dbv, dbh, dwhv, learning_rate=params.learning_rate / data.shape[0],
|
grad = self.entity.grad_compute(dbv, dbh, dwhv)
|
||||||
momentum=params.momentum, weight_decay=params.weight_decay)
|
|
||||||
# Adjust weights
|
# Adjust weights
|
||||||
self.entity.state_adjust(grad)
|
self.entity.state_adjust(grad)
|
||||||
|
|
||||||
|
|||||||
+2
-2
@@ -225,8 +225,8 @@ def train(entity: Entity, batch: Mat, status: Status):
|
|||||||
dwhv, dbv, dbh = cd_func(entity, mini_batch)
|
dwhv, dbv, dbh = cd_func(entity, mini_batch)
|
||||||
|
|
||||||
# Adjust weight and biases
|
# 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)
|
grad = entity.grad_compute(dbv, dbh, dwhv)
|
||||||
entity.state_adjust(grad)
|
entity.state_adjust(grad, 1.0/mini_batch.shape[0])
|
||||||
|
|
||||||
# check if status update is needed
|
# check if status update is needed
|
||||||
if status.want_report(round(progress)):
|
if status.want_report(round(progress)):
|
||||||
|
|||||||
@@ -40,7 +40,7 @@ def linear():
|
|||||||
model.init(0.1)
|
model.init(0.1)
|
||||||
|
|
||||||
# Load weights (if exists)
|
# Load weights (if exists)
|
||||||
model.load()
|
# model.load()
|
||||||
|
|
||||||
# Prepare training data
|
# Prepare training data
|
||||||
training_batch = np.random.randn(N_CASES, N_VIS, dtype=np.float64)
|
training_batch = np.random.randn(N_CASES, N_VIS, dtype=np.float64)
|
||||||
|
|||||||
@@ -5,14 +5,16 @@ from rbm.model import Model
|
|||||||
from rbm.entity import Entity, EntityParams, TrainingParams
|
from rbm.entity import Entity, EntityParams, TrainingParams
|
||||||
from rbm.matrix import Mat, np, read_armadillo
|
from rbm.matrix import Mat, np, read_armadillo
|
||||||
|
|
||||||
|
DO_HIDDEN_GAUSSIAN = True
|
||||||
|
|
||||||
class TestModel(Model):
|
class TestModel(Model):
|
||||||
def __init__(self, name: str, work_dir: str = '.', do_gaussian_hidden=False):
|
def __init__(self, name: str, work_dir: str = '.'):
|
||||||
super().__init__(name, work_dir)
|
super().__init__(name, work_dir)
|
||||||
|
|
||||||
if do_gaussian_hidden:
|
if DO_HIDDEN_GAUSSIAN:
|
||||||
# Hidden gaussian
|
# Hidden gaussian
|
||||||
self.unit1 = Entity((96*96, 16), EntityParams(do_gaussian_visible=True, do_gaussian_hidden=True),
|
self.unit1 = Entity((96*96, 16), EntityParams(do_gaussian_visible=True, do_gaussian_hidden=True),
|
||||||
TrainingParams(learning_rate=0.000001, momentum=0.9, num_epochs=1000))
|
TrainingParams(learning_rate=0.00001, momentum=0.9, num_epochs=1000))
|
||||||
else:
|
else:
|
||||||
# Hidden binary
|
# Hidden binary
|
||||||
self.unit1 = Entity((96 * 96, 333), EntityParams(do_gaussian_visible=True, do_gaussian_hidden=False),
|
self.unit1 = Entity((96 * 96, 333), EntityParams(do_gaussian_visible=True, do_gaussian_hidden=False),
|
||||||
@@ -35,10 +37,13 @@ if __name__ == "__main__":
|
|||||||
model = TestModel(prj_name, "results")
|
model = TestModel(prj_name, "results")
|
||||||
|
|
||||||
# Init state
|
# Init state
|
||||||
|
if DO_HIDDEN_GAUSSIAN:
|
||||||
|
model.init(0.01)
|
||||||
|
else:
|
||||||
model.init(0.1)
|
model.init(0.1)
|
||||||
|
|
||||||
# load state
|
# load state
|
||||||
model.load()
|
# model.load()
|
||||||
|
|
||||||
# Load train data
|
# Load train data
|
||||||
train_batch = read_armadillo(os.path.join(prj_root, f"{prj_name}.training.dat"))
|
train_batch = read_armadillo(os.path.join(prj_root, f"{prj_name}.training.dat"))
|
||||||
|
|||||||
Reference in New Issue
Block a user