increased difficulty
This commit is contained in:
@@ -24,11 +24,11 @@ from rbm.train import train
|
|||||||
from rbm.status import Status
|
from rbm.status import Status
|
||||||
|
|
||||||
# ── Hyper-parameters ──────────────────────────────────────────────────────────
|
# ── Hyper-parameters ──────────────────────────────────────────────────────────
|
||||||
TEXT = " JAY IS COOL" # the character sequence to learn (matches diagram example)
|
TEXT = " HALLO SUPER JENS. SUPER! " # the character sequence to learn (matches diagram example)
|
||||||
WIN = 3 # sliding-window width (= N in the diagram)
|
WIN = 3 # sliding-window width (= N in the diagram)
|
||||||
STRIDE = 1 # sliding-window step size
|
STRIDE = 1 # sliding-window step size
|
||||||
UNITS = 1
|
UNITS = 1
|
||||||
H_SIZE = 64 # hidden units per RBM cell
|
H_SIZE = 120 # hidden units per RBM cell
|
||||||
|
|
||||||
TRAIN_PARAMS = TrainingParams(
|
TRAIN_PARAMS = TrainingParams(
|
||||||
learning_rate = 0.05,
|
learning_rate = 0.05,
|
||||||
@@ -71,7 +71,7 @@ class RnnModel(Model):
|
|||||||
|
|
||||||
self.units: list[Entity] = []
|
self.units: list[Entity] = []
|
||||||
for _ in range(UNITS):
|
for _ in range(UNITS):
|
||||||
unit = Entity((WIN*vocab_size() + H_SIZE, H_SIZE), EntityParams(do_gaussian_visible=False, do_gaussian_hidden=False), TrainingParams(learning_rate=0.001, momentum=0.9, num_epochs=1000), enable_training=True)
|
unit = Entity((WIN*vocab_size() + H_SIZE, H_SIZE), EntityParams(do_gaussian_visible=False, do_gaussian_hidden=False), training_params=TRAIN_PARAMS)
|
||||||
self.units.append(unit)
|
self.units.append(unit)
|
||||||
|
|
||||||
def train(self, vc: Mat, status: Status = None):
|
def train(self, vc: Mat, status: Status = None):
|
||||||
@@ -91,17 +91,21 @@ class RnnModel(Model):
|
|||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
vec = idx2vec(ch2idx("JENS"))
|
if 1:
|
||||||
vec = clamp(vec, axis=1)
|
# Test of helper function
|
||||||
indices = vec2idx(vec, axis=1)
|
# ToDo: move to tests
|
||||||
ch_str = idx2ch(indices)
|
vec = idx2vec(ch2idx("JENS"))
|
||||||
print(ch_str)
|
vec = clamp(vec, axis=1)
|
||||||
|
indices = vec2idx(vec, axis=1)
|
||||||
|
ch_str = idx2ch(indices)
|
||||||
|
print(ch_str)
|
||||||
|
|
||||||
win_text = to_window(TEXT)
|
win_text = to_window(TEXT)
|
||||||
print(win_text)
|
print(win_text)
|
||||||
|
|
||||||
|
batch = to_batch(win_text)
|
||||||
|
print(batch)
|
||||||
|
|
||||||
batch = to_batch(win_text)
|
|
||||||
print(batch)
|
|
||||||
|
|
||||||
model = RnnModel(name='JayRnn', work_dir='results')
|
model = RnnModel(name='JayRnn', work_dir='results')
|
||||||
model.init(0.01)
|
model.init(0.01)
|
||||||
@@ -111,14 +115,14 @@ if __name__ == "__main__":
|
|||||||
# context will be updated after training
|
# context will be updated after training
|
||||||
c_train = np.zeros([len(batch), H_SIZE])
|
c_train = np.zeros([len(batch), H_SIZE])
|
||||||
vc_train = concat(batch, c_train, axis=1)
|
vc_train = concat(batch, c_train, axis=1)
|
||||||
# model.train(vc_train, status=Status())
|
model.train(vc_train, status=Status())
|
||||||
model.save()
|
model.save()
|
||||||
|
|
||||||
# test the model
|
# test the model
|
||||||
seed = str2vec('JA^')
|
seed = str2vec('HA^')
|
||||||
v_forward = seed.flatten()
|
v_forward = seed.flatten()
|
||||||
context = c_train[0]
|
context = c_train[0]
|
||||||
for i in range(len(TEXT)):
|
for i in range(len(TEXT)+5):
|
||||||
print(f"Forward {i:02d}: {vec2str(v_forward.reshape([WIN, vocab_size()]), axis=1)}")
|
print(f"Forward {i:02d}: {vec2str(v_forward.reshape([WIN, vocab_size()]), axis=1)}")
|
||||||
vc_forward = concat(v_forward.reshape([1, WIN*vocab_size()]), context.reshape([1, H_SIZE]), axis=1)
|
vc_forward = concat(v_forward.reshape([1, WIN*vocab_size()]), context.reshape([1, H_SIZE]), axis=1)
|
||||||
vc_predict = model.forward_step(vc_forward)
|
vc_predict = model.forward_step(vc_forward)
|
||||||
|
|||||||
Reference in New Issue
Block a user