"""Tests for StackRnn — shared-weights and unrolled (own-weights) modes.""" import numpy as _np_cpu from stack.rnn import StackRnn from rbm.matrix import np, convert from rbm.entity import EntityParams, TrainingParams from rbm.status import Status SENSORY_SIZE = 8 H_SIZE = 4 T = 16 # sequence length NUM_SEQ = 10 # sequences in the batch WORK_DIR = "../../results" _PARAMS = TrainingParams(learning_rate=0.05, momentum=0.5, num_epochs=5, do_rao_blackwell=True) def _make_sequences() -> _np_cpu.ndarray: rng = _np_cpu.random.RandomState(0) base = (rng.rand(NUM_SEQ, SENSORY_SIZE) > 0.5).astype(_np_cpu.float64) return _np_cpu.stack([base] * T, axis=1) # (NUM_SEQ, T, SENSORY_SIZE) def test_rnn_shared(): """Shared-weights mode: one entity reused at every time step.""" seqs = _make_sequences() rnn = StackRnn("test_rnn_shared", WORK_DIR) rnn.append(StackRnn.make_layer("layer0", SENSORY_SIZE, H_SIZE, EntityParams(), _PARAMS)) rnn.state_init(0.01) assert rnn.is_shared assert rnn.sensory_size() == SENSORY_SIZE assert rnn.h_size() == H_SIZE rnn.train(seqs) rnn.reset(batch_size=1) for t in range(T): h = rnn.step(np.array(seqs[0, t][None, :])) assert h.shape == (1, H_SIZE), f"bad shape at t={t}" recon = rnn.reconstruct(h) assert recon.shape == (1, SENSORY_SIZE) print(f"Original : {convert(np.array(seqs[0, -1][None, :]))}") print(f"Recon : {convert(recon)}") print("test_rnn_shared: [passed]") def test_rnn_unrolled(): """Unrolled mode: T entities, one per time step, each with own weights.""" seqs = _make_sequences() rnn = StackRnn("test_rnn_unrolled", WORK_DIR) for layer in StackRnn.make_unrolled(T, SENSORY_SIZE, H_SIZE, EntityParams(), _PARAMS): rnn.append(layer) rnn.state_init(0.01) assert not rnn.is_shared assert rnn.num_layers() == T rnn.train(seqs) # Inference — _t advances through all T entities rnn.reset(batch_size=1) for t in range(T): h = rnn.step(np.array(seqs[0, t][None, :])) assert h.shape == (1, H_SIZE), f"bad shape at t={t}" assert rnn._t == t + 1 recon = rnn.reconstruct(h) assert recon.shape == (1, SENSORY_SIZE) # next_entity wraps around modularly rnn._t = T + 3 assert rnn.next_entity() is rnn.from_index(3).entity print("test_rnn_unrolled: [passed]") def test_rnn_save_load(): seqs = _make_sequences() rnn = StackRnn("test_rnn_sl", WORK_DIR) for layer in StackRnn.make_unrolled(T, SENSORY_SIZE, H_SIZE, EntityParams(), _PARAMS): rnn.append(layer) rnn.state_init(0.01) rnn.train(seqs) rnn.state_save() rnn.state_load() print("test_rnn_save_load: [passed]") if __name__ == "__main__": test_rnn_shared() test_rnn_unrolled() test_rnn_save_load() print("All RNN tests passed.")