177 lines
5.0 KiB
Python
177 lines
5.0 KiB
Python
#!/usr/bin/env python3
|
|
"""Character-level RNN-RBM (unrolled) — see docs/Rnn.drawio.png.
|
|
|
|
Diagram recap
|
|
t=0: v[0]={0} v[1:N]=[' ',' ','J'] W[0] → h
|
|
t=1: v[0]=h₀ v[1:N]=[' ','J','A'] W[1] → h
|
|
...
|
|
t=M-1: v[0]=h_{M-2} v[1:N]=['J','A','Y'] W[M-1] → h
|
|
|
|
v[0] = recurrent context (previous hidden state, or zeros at t=0)
|
|
v[1:N] = N_WIN one-hot-encoded characters concatenated (the sliding window)
|
|
W[t] = RBM weight matrix for time step t (one per step → unrolled mode)
|
|
"""
|
|
import sys
|
|
import os
|
|
|
|
sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'src'))
|
|
|
|
from rbm.entity import EntityParams, TrainingParams, Entity
|
|
from rbm.matrix import np, Mat
|
|
from model.model import Model
|
|
from stack.rnn_helper import vocab_size, shift_left, clamp, concat, split, vec2idx, idx2ch, ch2idx, idx2vec
|
|
from rbm.train import train
|
|
from rbm.status import Status
|
|
|
|
# ── Hyper-parameters ──────────────────────────────────────────────────────────
|
|
TEXT = "Call me Ishmael. Some years ago" # the character sequence to learn (matches diagram example)
|
|
WIN = 2 # sliding-window width (= N in the diagram)
|
|
STRIDE = 1 # sliding-window step size
|
|
UNITS = 3
|
|
H_SIZE = 128 # hidden units per RBM cell
|
|
NUM_EPOCHS = 1000
|
|
NUM_ITERATIONS = 1
|
|
|
|
TRAIN_PARAMS = TrainingParams(
|
|
learning_rate = 0.1,
|
|
momentum = 0.5,
|
|
num_epochs = NUM_EPOCHS,
|
|
do_rao_blackwell = True,
|
|
num_gibbs_samples= 1
|
|
)
|
|
#WIN=3
|
|
#STRIDE=1
|
|
#UNIT=3
|
|
#WIN | |
|
|
#U0: "JAY IS COOL"
|
|
#U1: "AY IS COOL "
|
|
#U2: "Y IS COOL "
|
|
|
|
def vec2str(mat: Mat, axis=0):
|
|
return idx2ch(vec2idx(mat, axis))
|
|
|
|
def str2vec(ch_str: str) -> Mat:
|
|
return idx2vec(ch2idx(ch_str))
|
|
|
|
def to_window(text: str) -> list[str]:
|
|
result = []
|
|
_win = [' ']*UNITS
|
|
for ch in text:
|
|
for _i in range(WIN-1):
|
|
_win = _win[1:WIN] + [ch]
|
|
result += _win
|
|
|
|
return result
|
|
|
|
def to_batch(_win_text: list[str], adv=1):
|
|
_batch = np.zeros([len(_win_text), WIN*vocab_size()])
|
|
for _i in range(len(TEXT)-UNITS):
|
|
win_str = ''
|
|
for _j in range(WIN):
|
|
win_str += _win_text[(_i+adv)*WIN+_j]
|
|
_batch[_i, :] = str2vec(win_str).flatten()
|
|
return _batch
|
|
|
|
def to_batch_(_win_str: list[str], delay=0):
|
|
_batch = np.zeros([len(_win_str), WIN, vocab_size()])
|
|
for i in range(len(TEXT)):
|
|
for j in range(WIN):
|
|
_vec = str2vec(_win_str[i+j])
|
|
_batch[i, j, :] = _vec
|
|
return _batch
|
|
|
|
def batch_delay(_batch: Mat, delay=0):
|
|
result = _batch
|
|
if delay > 0:
|
|
result[0:-delay] = _batch[delay:]
|
|
result[-delay:] = np.zeros([delay, _batch.shape[1]])
|
|
return result
|
|
|
|
def vc2char(_vc: Mat):
|
|
_v, _ = split(_vc.reshape(1, WIN*vocab_size()+H_SIZE), H_SIZE, axis=1)
|
|
return v2char(_v)
|
|
|
|
def v2char(_v: Mat):
|
|
return vec2str(_v.reshape([WIN, vocab_size()]), axis=1)
|
|
|
|
class RnnModel(Model):
|
|
def __init__(self, name: str, work_dir: str = '.'):
|
|
super().__init__(name, work_dir)
|
|
|
|
self.units: list[Entity] = []
|
|
for index in range(UNITS):
|
|
unit = Entity((WIN*vocab_size() + H_SIZE, H_SIZE), EntityParams(do_gaussian_visible=False, do_gaussian_hidden=False), training_params=TRAIN_PARAMS, index=index)
|
|
self.units.append(unit)
|
|
|
|
def seq_len(self):
|
|
return len(self.units)
|
|
|
|
def train(self, _win_text: list[str], status: Status = None):
|
|
_c = np.zeros([len(_win_text), H_SIZE])
|
|
for index, unit in enumerate(self.units):
|
|
_batch = to_batch(_win_text, index)
|
|
_vc = concat(_batch, _c, axis=1)
|
|
for i in range(_vc.shape[0]):
|
|
print(f"train: {unit.name}:{vc2char(_vc[i,:])}")
|
|
train(unit, _vc, status)
|
|
# For the next unit: Update context portion of vc
|
|
_c = unit.forward(_vc)
|
|
|
|
|
|
def forward_step(self, _vc: Mat):
|
|
text = ''
|
|
for unit in self.units:
|
|
print(f"forward_step in : {unit.name}:{vc2char(_vc)}")
|
|
_c = unit.forward(_vc)
|
|
_vc = unit.reconstruct(_c)
|
|
print(f"forward_step out: {unit.name}:{vc2char(_vc)}")
|
|
text += vc2char(_vc)[WIN-1]
|
|
_v, _ = split(_vc, H_SIZE, axis=1)
|
|
_v_cl = clamp(_v.reshape([WIN, vocab_size()]), axis=1)
|
|
_v = shift_left(_v_cl.reshape([1, WIN*vocab_size()]), vocab_size())
|
|
_vc = concat(_v, _c, axis=1)
|
|
return _vc, text
|
|
|
|
def forward(self, x: Mat):
|
|
pass
|
|
|
|
|
|
if __name__ == "__main__":
|
|
if 1:
|
|
# Test of helper function
|
|
# ToDo: move to tests
|
|
vec = idx2vec(ch2idx("JENS"))
|
|
vec = clamp(vec, axis=1)
|
|
indices = vec2idx(vec, axis=1)
|
|
ch_str = idx2ch(indices)
|
|
print(ch_str)
|
|
|
|
win_text = to_window(TEXT)
|
|
print(win_text)
|
|
|
|
batch = to_batch(win_text)
|
|
print(batch)
|
|
|
|
|
|
model = RnnModel(name='JayRnn', work_dir='results')
|
|
model.init(0.01)
|
|
model.load()
|
|
|
|
# vc contains vis + context
|
|
# context will be updated after training
|
|
model.train(win_text, status=Status())
|
|
model.save()
|
|
|
|
# test the model
|
|
seed_str = ' '
|
|
seed_padded = seed_str + ' '*(WIN-len(seed_str))
|
|
v_test = str2vec(seed_padded).reshape([1, WIN*vocab_size()])
|
|
c_test = np.zeros([1, H_SIZE])
|
|
vc_test = concat(v_test, c_test, axis=1)
|
|
text = ''
|
|
for i in range(len(TEXT)):
|
|
vc_test, text_frag = model.forward_step(vc_test)
|
|
text += text_frag
|
|
|
|
print(text)
|