[add] Rnn.py: unrolled character-level RNN-RBM matching docs/Rnn.drawio.png

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-06-03 10:49:46 +02:00
co-authored by Claude Sonnet 4.6
parent 3edc124a5c
commit df4a9b2ac2
3 changed files with 234 additions and 0 deletions
+122
View File
@@ -0,0 +1,122 @@
#!/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'))
import numpy as np
from stack.rnn import StackRnn
from rbm.entity import EntityParams, TrainingParams
from rbm.matrix import convert
# ── Hyper-parameters ──────────────────────────────────────────────────────────
TEXT = " JAY" # the character sequence to learn (matches diagram example)
N_WIN = 3 # sliding-window width (= N in the diagram)
H_SIZE = 8 # hidden units per RBM cell
TRAIN_PARAMS = TrainingParams(
learning_rate = 0.05,
momentum = 0.9,
num_epochs = 1000,
do_rao_blackwell = True,
)
N_REPS = 300 # training repetitions of the sequence
# ── Vocabulary ────────────────────────────────────────────────────────────────
chars = sorted(set(TEXT))
VOCAB_SIZE = len(chars)
c2i = {c: i for i, c in enumerate(chars)}
i2c = {i: c for c, i in c2i.items()}
SENS_SIZE = N_WIN * VOCAB_SIZE # width of v[1:N]
print(f"Vocab ({VOCAB_SIZE}): {chars}")
print(f"N_WIN={N_WIN} H_SIZE={H_SIZE} SENS_SIZE={SENS_SIZE}")
# ── Encoding helpers ──────────────────────────────────────────────────────────
def encode_window(window: str) -> np.ndarray:
"""Encode N_WIN characters as a flat one-hot vector (v[1:N])."""
vec = np.zeros(SENS_SIZE, dtype=np.float64)
for j, c in enumerate(window):
vec[j * VOCAB_SIZE + c2i[c]] = 1.0
return vec
def decode_char(x: np.ndarray, position: int) -> str:
"""Decode the one-hot slot at `position` inside a flat window vector."""
seg = x[position * VOCAB_SIZE : (position + 1) * VOCAB_SIZE]
return i2c[int(np.argmax(seg))]
# ── Build training sequences (num_seq, T, SENS_SIZE) ─────────────────────────
# M = number of sliding windows per period of TEXT.
M = len(TEXT) - N_WIN + 1 # time steps per sequence
seqs = np.zeros((N_REPS, M, SENS_SIZE), dtype=np.float64)
for s in range(N_REPS):
for t in range(M):
seqs[s, t] = encode_window(TEXT[t : t + N_WIN])
print(f"\nSequences: {N_REPS} × {M} steps (T={M})")
# ── Build unrolled StackRnn: M layers, one W[t] per time step ─────────────────
rnn = StackRnn("rnn_char", "results/rnn_char")
for layer in StackRnn.make_unrolled(M, SENS_SIZE, H_SIZE, EntityParams(), TRAIN_PARAMS):
rnn.append(layer)
rnn.state_init(0.01)
print(f"Training…")
rnn.train(seqs)
rnn.state_save()
print("Weights saved to results/rnn_char/")
# ── Generate text ─────────────────────────────────────────────────────────────
def generate(seed: str, num_chars: int = 20) -> str:
"""Generate text by repeatedly reconstructing the next character.
The model reconstructs v[1:N] from the hidden state h. We take the last
one-hot slot of the reconstructed window as the next predicted character,
then slide the window by one.
"""
assert len(seed) >= N_WIN, f"seed must be ≥ {N_WIN} chars"
rnn.reset(batch_size=1)
# Prime: step through all but the last window of the seed
for t in range(len(seed) - N_WIN):
rnn.step(encode_window(seed[t : t + N_WIN])[None, :])
window = seed[-N_WIN:]
result = list(seed)
for _ in range(num_chars):
h = rnn.step(encode_window(window)[None, :])
x_recon = convert(rnn.reconstruct(h))[0]
next_char = decode_char(x_recon, N_WIN - 1)
result.append(next_char)
window = window[1:] + next_char
return ''.join(result)
seed = TEXT[:N_WIN]
print(f"\nGenerated (seed={seed!r}):")
print(repr(generate(seed, num_chars=20)))
def main():
_seed = TEXT[:N_WIN]
print(f"\nGenerated (seed={_seed!r}):")
print(repr(generate(_seed, num_chars=20)))
if __name__ == "__main__":
main()