File size: 4,845 Bytes
b4e986d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
"""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