diff --git a/JayRnn.py b/JayRnn.py new file mode 100644 index 0000000..af18c30 --- /dev/null +++ b/JayRnn.py @@ -0,0 +1,113 @@ +#!/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, concat, vec2idx, idx2ch, ch2idx, idx2vec +from rbm.train import train +from rbm.status import Status + +# ── Hyper-parameters ────────────────────────────────────────────────────────── +TEXT = "JAY IS COOL" # the character sequence to learn (matches diagram example) +WIN = 3 # sliding-window width (= N in the diagram) +STRIDE = 1 # sliding-window step size +UNITS = 1 +H_SIZE = 64 # hidden units per RBM cell + +TRAIN_PARAMS = TrainingParams( + learning_rate = 0.05, + momentum = 0.9, + num_epochs = 1000, + do_rao_blackwell = True, +) +#WIN=3 +#STRIDE=1 +#UNIT=3 +#WIN | | +#U0: "JAY IS COOL" +#U1: "AY IS COOL " +#U2: "Y IS COOL " + +def vec2str(mat: Mat): + result = '' + for vec in mat: + result += idx2ch(vec2idx(vec)) + return result + +def str2vec(ch_str: str) -> Mat: + result = Mat([len(ch_str), vocab_size()]) + for i, ch in enumerate(ch_str): + vec = idx2vec(ch2idx(ch)) + result[i, :] = vec + return result + +def to_window(text: str): + win_text = [] + text_padded = text + ' '*WIN + for i in range(len(TEXT)): + win_text.append(text_padded[i:i+WIN]) + + return win_text + +def to_batch(_win_str: list[str]): + _batch = np.zeros([len(TEXT), WIN*vocab_size()]) + for i, win in enumerate(_win_str): + _batch[i, :] = str2vec(win).flatten() + return _batch + + +class RnnModel(Model): + def __init__(self, name: str, work_dir: str = '.'): + super().__init__(name, work_dir) + + self.units: list[Entity] = [] + 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) + self.units.append(unit) + + def train(self, batch: Mat, status: Status = None): + c = np.zeros([len(batch), H_SIZE]) + for i, unit in enumerate(self.units): + vc = concat(batch, c) + train(unit, vc, status) + c = unit.forward(vc) + + def forward_step(self, v_curr: Mat): + c = np.zeros([1, H_SIZE]) + z = np.zeros(v_curr.shape) + for i, unit in enumerate(self.units): + pass + + def forward(self, x: Mat): + pass + + +if __name__ == "__main__": + 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() + model.train(batch, status=Status()) + model.save() + diff --git a/docs/Rnn.drawio b/docs/Rnn.drawio index cfe0630..f1b71ef 100644 --- a/docs/Rnn.drawio +++ b/docs/Rnn.drawio @@ -1,23 +1,23 @@ - + - + - + - + - + - + @@ -29,28 +29,28 @@ - + - + - + - + - + - + - + @@ -62,28 +62,28 @@ - + - + - + - + - + - + - + @@ -95,16 +95,22 @@ - + - + - + + + + + + +