palimpseste-max / tests /test_memory.py
thefinalboss's picture
Upload tests/test_memory.py with huggingface_hub
b4e986d verified
Raw
History Blame Contribute Delete
4.85 kB
"""Tests for the append-only knowledge base ``M`` (``palimseste.memory``).
Core guarantees under test:
- append-only: nothing is ever deleted; len grows monotonically
- O(1) write (insert amortized cost, no retraining)
- soft decay never reaches zero (dormant memories can re-awaken)
- meta subspace is disjoint from the main recall space
- LSH candidates are returned for stored addresses
"""
from __future__ import annotations
import math
import numpy as np
import pytest
from palimseste import hv
from palimseste.memory import Memory, Trace
def _rand_mem(D=1000, **kw) -> Memory:
return Memory(D=D, rng=np.random.default_rng(42), **kw)
def test_write_returns_trace_and_grows():
mem = _rand_mem()
rng = np.random.default_rng(1)
assert len(mem) == 0
t0 = mem.write(hv.random_hv(D=1000, rng=rng), hv.random_hv(D=1000, rng=rng))
assert isinstance(t0, Trace)
assert t0.id == 0
assert len(mem) == 1
t1 = mem.write(hv.random_hv(D=1000, rng=rng), hv.random_hv(D=1000, rng=rng))
assert t1.id == 1
assert len(mem) == 2
def test_write_dimension_mismatch():
mem = _rand_mem(D=500)
a = hv.random_hv(D=500)
v = hv.random_hv(D=999)
with pytest.raises(ValueError):
mem.write(a, v)
with pytest.raises(ValueError):
mem.write(v, a)
def test_append_only_never_deletes():
mem = _rand_mem()
rng = np.random.default_rng(2)
ids_before = []
for _ in range(50):
t = mem.write(hv.random_hv(D=1000, rng=rng), hv.random_hv(D=1000, rng=rng))
ids_before.append(t.id)
# writing more must not remove any prior trace
for _ in range(50):
mem.write(hv.random_hv(D=1000, rng=rng), hv.random_hv(D=1000, rng=rng))
ids_after = [t.id for t in mem.traces]
assert ids_after[:50] == ids_before
assert len(mem) == 100
def test_weight_never_zero():
mem = _rand_mem(D=500, decay={"half_life": 0.01, "floor": 1e-3})
rng = np.random.default_rng(3)
a = hv.random_hv(D=500, rng=rng)
v = hv.random_hv(D=500, rng=rng)
tr = mem.write(a, v, weight=1.0)
# wait a bit so decay kicks in
import time
time.sleep(0.05)
w = mem.current_weight(tr, now=time.monotonic())
assert w > 0.0 # never zero
assert w < 1.0 # decayed
def test_no_decay_keeps_weight_constant():
mem = _rand_mem(D=500, decay={"half_life": math.inf, "floor": 1e-3})
rng = np.random.default_rng(4)
tr = mem.write(hv.random_hv(D=500, rng=rng), hv.random_hv(D=500, rng=rng), weight=0.7)
assert mem.current_weight(tr) == pytest.approx(0.7)
def test_meta_subspace_disjoint():
mem = _rand_mem(D=1000)
rng = np.random.default_rng(5)
# write a normal and a meta trace
normal = mem.write(hv.random_hv(D=1000, rng=rng), hv.random_hv(D=1000, rng=rng))
meta = mem.write(
hv.random_hv(D=1000, rng=rng),
hv.random_hv(D=1000, rng=rng),
meta=True,
tag="kernel_radius",
)
assert normal.meta is False
assert meta.meta is True
assert mem.meta_traces[0].id == meta.id
assert len(mem.traces) == 1
assert len(mem.meta_traces) == 1
def test_candidates_returned_for_stored_address():
mem = _rand_mem(D=1000)
rng = np.random.default_rng(6)
addrs = [hv.random_hv(D=1000, rng=rng) for _ in range(20)]
for a in addrs:
mem.write(a, hv.random_hv(D=1000, rng=rng))
# querying a stored address should include itself in candidates
cand = mem.candidates(addrs[5])
assert 5 in cand
def test_trace_rejects_zero_weight():
with pytest.raises(ValueError):
Trace(id=0, address=hv.random_hv(D=10), value=hv.random_hv(D=10), weight=0.0)
with pytest.raises(ValueError):
Trace(id=0, address=hv.random_hv(D=10), value=hv.random_hv(D=10), weight=-1.0)
def test_stats_sensible():
mem = _rand_mem(D=500)
rng = np.random.default_rng(7)
for _ in range(30):
mem.write(hv.random_hv(D=500, rng=rng), hv.random_hv(D=500, rng=rng), weight=1.0)
s = mem.stats()
assert s.n_traces == 30
assert s.n_meta == 0
assert s.mean_weight == pytest.approx(1.0)
assert s.min_weight == pytest.approx(1.0)
assert s.lsh_size == 30
def test_rebuild_index_preserves_traces():
mem = _rand_mem(D=500)
rng = np.random.default_rng(8)
for _ in range(40):
mem.write(hv.random_hv(D=500, rng=rng), hv.random_hv(D=500, rng=rng))
n_before = len(mem)
mem.rebuild_index()
assert len(mem) == n_before
assert mem.index.size == n_before
def test_get_meta():
mem = _rand_mem(D=500)
rng = np.random.default_rng(9)
a = hv.random_hv(D=500, rng=rng)
v = hv.random_hv(D=500, rng=rng)
tr = mem.write(a, v, meta=True, tag="test")
got = mem.get_meta(tr.id)
assert got is not None
assert got.id == tr.id
assert mem.get_meta(99999999) is None