File size: 5,699 Bytes
335815c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
#!/usr/bin/env python3
"""
example.py — Usage demonstrations for hv_acausal_memory.

NumPy only.
"""

import numpy as np

from hv_acausal_memory import (
    HyperNSV, sparse_random_ternary, add_noise, make_partial_query,
)


def demo_store_retrieve():
    print("=" * 72)
    print("Demo 1 — store and retrieve")
    print("=" * 72)
    print()

    rng = np.random.default_rng(0)
    mem = HyperNSV(D=2000, nnz=80, max_anchors=100, seed=0)

    for i in range(50):
        key = sparse_random_ternary(2000, 80, rng)
        val = sparse_random_ternary(2000, 80, rng)
        mem.store(key, val)

    print(f"  stored {mem.n_active} anchors")

    # exact query
    query = mem.keys[7].copy()
    results = mem.retrieve(query, top_k=3)
    print(f"  exact query for anchor 7:")
    for idx, score, sim in results:
        print(f"    anchor {idx}:  score {score:.4f}  sim {sim:.4f}")
    print()

    # noisy query
    noisy = add_noise(mem.keys[7].copy(), 0.2, rng)
    results = mem.retrieve(noisy, top_k=3)
    print(f"  noisy query (p=0.2):")
    for idx, score, sim in results:
        print(f"    anchor {idx}:  score {score:.4f}  sim {sim:.4f}")
    print()


def demo_forget_gate():
    print("=" * 72)
    print("Demo 2 — forget-gate under pressure")
    print("=" * 72)
    print()

    for n_store, max_c in [(10, 100), (90, 100)]:
        rng = np.random.default_rng(0)
        mem = HyperNSV(D=500, nnz=30, max_anchors=max_c, seed=0)
        for _ in range(n_store):
            mem.store(sparse_random_ternary(500, 30, rng),
                      sparse_random_ternary(500, 30, rng))
        pressure = n_store / max_c
        beta = mem.beta_base + (mem.beta_max - mem.beta_base) * pressure
        print(f"  {n_store}/{max_c}  pressure {pressure:.2f}  β {beta:.4f}")
        for horizon in [0, 10, 30, 100]:
            w = float(mem.weights[:n_store].mean())
            print(f"    t={horizon:>3d}  mean weight {w:.4f}")
            for _ in range(10 if horizon == 0 else
                           (10 if horizon == 10 else
                            (20 if horizon == 30 else 70))):
                mem.step_time(dt=1.0)
        print()


def demo_resurrection():
    print("=" * 72)
    print("Demo 3 — resurrection")
    print("=" * 72)
    print()

    rng = np.random.default_rng(0)
    mem = HyperNSV(D=1000, nnz=40, max_anchors=100, seed=0)
    for _ in range(50):
        mem.store(sparse_random_ternary(1000, 40, rng),
                  sparse_random_ternary(1000, 40, rng))

    # kill all via time
    for _ in range(200):
        mem.step_time(dt=1.0)

    print(f"  after 200 steps: "
          f"{mem.stats()['n_live']} live, "
          f"{mem.stats()['n_dead']} dead")

    # revive 10 with exact keys
    dead_idx = np.where(mem.weights[:50] <= mem.dead_threshold)[0]
    picked = rng.choice(dead_idx, min(10, len(dead_idx)), replace=False)
    revived_count = 0
    for idx in picked:
        revived_count += len(mem.resurrect(mem.keys[idx].copy(),
                                            threshold=0.3, boost=0.5))
    print(f"  exact-key reminders: {revived_count}/10 revived")
    print(f"  after revival: "
          f"{mem.stats()['n_live']} live, "
          f"{mem.stats()['n_dead']} dead")
    print()


def demo_acausal_prefetch():
    print("=" * 72)
    print("Demo 4 — acausal prefetch from partial query")
    print("=" * 72)
    print()

    rng = np.random.default_rng(0)
    mem = HyperNSV(D=2000, nnz=80, max_anchors=100, seed=0)
    for _ in range(50):
        mem.store(sparse_random_ternary(2000, 80, rng),
                  sparse_random_ternary(2000, 80, rng))

    for keep in [0.7, 0.5, 0.3]:
        partial = make_partial_query(mem.keys[13].copy(), keep,
                                      0.0, rng=rng)
        results = mem.acausal_prefetch(partial, top_k=3)
        target_in_top3 = any(r[0] == 13 for r in results)
        top1 = results[0][0] if results else -1
        print(f"  keep={keep:.1f}  top-1={top1}  "
              f"target in top-3: {target_in_top3}")
    print()


def demo_capacity():
    print("=" * 72)
    print("Demo 5 — capacity at two noise levels")
    print("=" * 72)
    print()

    for noise in [0.0, 0.2]:
        print(f"  noise p={noise}:")
        for n in [100, 500, 2000]:
            rng = np.random.default_rng(0)
            mem = HyperNSV(D=2000, nnz=80, max_anchors=n + 10, seed=0)
            for _ in range(n):
                mem.store(sparse_random_ternary(2000, 80, rng),
                          sparse_random_ternary(2000, 80, rng))
            correct = 0
            for _ in range(50):
                target = int(rng.integers(0, n))
                query = add_noise(mem.keys[target].copy(), noise, rng)
                results = mem.retrieve(query, top_k=1)
                if results and results[0][0] == target:
                    correct += 1
            print(f"    N={n:>5d}  acc {correct/50:.3f}")
        print()


def demo_stats():
    print("=" * 72)
    print("Demo 6 — bank statistics")
    print("=" * 72)
    print()

    rng = np.random.default_rng(0)
    mem = HyperNSV(D=5000, nnz=200, max_anchors=500, seed=0)
    for _ in range(200):
        mem.store(sparse_random_ternary(5000, 200, rng),
                  sparse_random_ternary(5000, 200, rng))

    stats = mem.stats()
    for k, v in stats.items():
        if isinstance(v, float):
            print(f"  {k:<22s} = {v:.4f}")
        else:
            print(f"  {k:<22s} = {v}")
    print()


def main():
    demo_store_retrieve()
    demo_forget_gate()
    demo_resurrection()
    demo_acausal_prefetch()
    demo_capacity()
    demo_stats()


if __name__ == "__main__":
    main()