refactored
This commit is contained in:
@@ -5,12 +5,14 @@ from rbm.label import Label
|
|||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
voc_size = 100
|
voc_size = 100
|
||||||
do_train = False
|
do_train = True
|
||||||
work_dir = "../../results"
|
work_dir = "../../results"
|
||||||
prj_root = "/home/jens/work/repos/Rbm"
|
prj_root = "/home/jens/work/repos/Rbm"
|
||||||
|
|
||||||
# Create Labeler
|
# Create Labeler
|
||||||
enc = Label(voc_size, 16, Label.EncodingType.OneHot, work_dir)
|
label_w = 5
|
||||||
|
label_h = 5
|
||||||
|
enc = Label(voc_size, label_w*label_h, Label.EncodingType.Binary, work_dir)
|
||||||
|
|
||||||
# Load
|
# Load
|
||||||
enc.fitter.load()
|
enc.fitter.load()
|
||||||
@@ -34,7 +36,7 @@ if __name__ == "__main__":
|
|||||||
|
|
||||||
fig, axes = plt.subplots(1, len(encoded), figsize=(12, 3))
|
fig, axes = plt.subplots(1, len(encoded), figsize=(12, 3))
|
||||||
for index, inp in enumerate(encoded):
|
for index, inp in enumerate(encoded):
|
||||||
img = np.reshape(inp, (4, 4))
|
img = np.reshape(inp, (label_w, label_h))
|
||||||
axes[index].imshow(img)
|
axes[index].imshow(img)
|
||||||
axes[index].axis('off')
|
axes[index].axis('off')
|
||||||
axes[index].set_title(f'{test_labels[index]}')
|
axes[index].set_title(f'{test_labels[index]}')
|
||||||
|
|||||||
@@ -7,9 +7,9 @@ from rbm.entity import Entity, EntityParams, TrainingParams
|
|||||||
from rbm.matrix import Mat, np
|
from rbm.matrix import Mat, np
|
||||||
|
|
||||||
class LabelLearner(Model):
|
class LabelLearner(Model):
|
||||||
def __init__(self, name: str, work_dir: str = '.'):
|
def __init__(self, name: str, dim, work_dir: str = '.'):
|
||||||
super().__init__(name, work_dir)
|
super().__init__(name, work_dir)
|
||||||
self.unit1 = Entity((16, 8), EntityParams(), TrainingParams(learning_rate=0.01, momentum=0.9, do_rao_blackwell=True, num_epochs=10000))
|
self.unit1 = Entity(dim, EntityParams(do_gaussian_hidden=False), TrainingParams(learning_rate=0.01, momentum=0.9, do_rao_blackwell=True, num_epochs=10000))
|
||||||
|
|
||||||
def forward(self, x: Mat):
|
def forward(self, x: Mat):
|
||||||
x = self.unit1.forward(x)
|
x = self.unit1.forward(x)
|
||||||
@@ -21,21 +21,30 @@ class LabelLearner(Model):
|
|||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
voc_size = 100
|
voc_size = 100
|
||||||
|
do_train = False
|
||||||
work_dir = "../../results"
|
work_dir = "../../results"
|
||||||
prj_root = "/home/jens/work/repos/Rbm"
|
prj_root = "/home/jens/work/repos/Rbm"
|
||||||
|
|
||||||
# Create Labeler
|
# Create Labeler
|
||||||
enc = Label(voc_size, 16, work_dir)
|
label_w = 5
|
||||||
|
label_h = 5
|
||||||
|
enc = Label(voc_size, label_w*label_h, Label.EncodingType.OneHot, work_dir)
|
||||||
|
|
||||||
# Load
|
# Load
|
||||||
enc.fitter.load()
|
enc.fitter.load()
|
||||||
|
|
||||||
|
# Train
|
||||||
|
if do_train:
|
||||||
|
train_labels = list(range(voc_size))
|
||||||
|
enc.fit(train_labels)
|
||||||
|
enc.fitter.save()
|
||||||
|
|
||||||
# Encode test labels
|
# Encode test labels
|
||||||
test_labels = np.array([1, 23, 99, 37, 55, 7, 31, 10, 19, 70])
|
test_labels = np.array([1, 23, 99, 37, 55, 7, 31, 10, 19, 70])
|
||||||
encoded, _ = enc.encode(test_labels)
|
encoded, _ = enc.encode(test_labels)
|
||||||
|
|
||||||
# Create model
|
# Create model
|
||||||
model = LabelLearner("Label-Learner", work_dir)
|
model = LabelLearner("Label-Learner", (label_w*label_h, 16), work_dir)
|
||||||
|
|
||||||
# Init state
|
# Init state
|
||||||
model.init(0.01)
|
model.init(0.01)
|
||||||
@@ -52,7 +61,7 @@ if __name__ == "__main__":
|
|||||||
# Plot
|
# Plot
|
||||||
fig, axes = plt.subplots(1, len(encoded), figsize=(12, 3))
|
fig, axes = plt.subplots(1, len(encoded), figsize=(12, 3))
|
||||||
for index, inp in enumerate(encoded):
|
for index, inp in enumerate(encoded):
|
||||||
img = np.reshape(inp, (4, 4))
|
img = np.reshape(inp, (label_w, label_h))
|
||||||
axes[index].imshow(img)
|
axes[index].imshow(img)
|
||||||
axes[index].axis('off')
|
axes[index].axis('off')
|
||||||
axes[index].set_title(f'{test_labels[index]}')
|
axes[index].set_title(f'{test_labels[index]}')
|
||||||
|
|||||||
Reference in New Issue
Block a user