Solving the Clive Wearing Problem: One-Shot Episodic Memory for Frozen Transformers
Tommi Niemi / Rotko Networks Hidden-state episodic memory for frozen transformers. No gradients. Teach via one forward pass, recall via cosine similarity + logit injection. 200-line Python reproduction included. pip install transformers torch numpy && python python/epimem.py
This commit is contained in:
2
.gitignore
vendored
Normal file
2
.gitignore
vendored
Normal file
@@ -0,0 +1,2 @@
|
|||||||
|
models/*.onnx
|
||||||
|
models/*.onnx.data
|
||||||
21
LICENSE
Normal file
21
LICENSE
Normal file
@@ -0,0 +1,21 @@
|
|||||||
|
MIT License
|
||||||
|
|
||||||
|
Copyright (c) 2026 Tommi Niemi / Rotko Networks
|
||||||
|
|
||||||
|
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||||
|
of this software and associated documentation files (the "Software"), to deal
|
||||||
|
in the Software without restriction, including without limitation the rights
|
||||||
|
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||||
|
copies of the Software, and to permit persons to whom the Software is
|
||||||
|
furnished to do so, subject to the following conditions:
|
||||||
|
|
||||||
|
The above copyright notice and this permission notice shall be included in all
|
||||||
|
copies or substantial portions of the Software.
|
||||||
|
|
||||||
|
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||||
|
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||||
|
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||||
|
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||||
|
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||||
|
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||||
|
SOFTWARE.
|
||||||
105
README.md
Normal file
105
README.md
Normal file
@@ -0,0 +1,105 @@
|
|||||||
|
<br />
|
||||||
|
<div align="center">
|
||||||
|
|
||||||
|
<h3 align="center">Solving the Clive Wearing Problem: One-Shot Episodic Memory for Frozen Transformers</h3>
|
||||||
|
|
||||||
|
<p align="center">
|
||||||
|
Gradient-free persistent learning through hidden-state episodic recall.
|
||||||
|
<br />
|
||||||
|
<a href="paper.md"><img src="https://img.shields.io/badge/Paper-Markdown-blue?style=flat-square" alt="Paper"></a>
|
||||||
|
</p>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Abstract
|
||||||
|
|
||||||
|
Teach a frozen language model new facts **without gradient descent**. Store the model's own hidden states as episodic memories. On recall, inject stored representations as logit biases. One-shot. Persistent. No weight modification.
|
||||||
|
|
||||||
|
## Key Result
|
||||||
|
|
||||||
|
A frozen Qwen 2.5 0.5B taught three facts about a fictional entity:
|
||||||
|
|
||||||
|
| Prompt | Taught | Recalled | Similarity |
|
||||||
|
|--------|--------|----------|:---:|
|
||||||
|
| "The capital of Zyphraxia is" | "Novaheim" | "Novaheim, a city of 100" | 1.000 |
|
||||||
|
| "The ruler of Zyphraxia is" | "Queen Stellara" | "Queen Stellara. She is a beautiful woman" | 1.000 |
|
||||||
|
| "The currency of Zyphraxia is" | "Glimmers" | "Glimmers. The currency is divided into" | 1.000 |
|
||||||
|
|
||||||
|
**No gradients computed at any point.**
|
||||||
|
|
||||||
|
## Quick Start
|
||||||
|
|
||||||
|
```bash
|
||||||
|
pip install transformers torch numpy
|
||||||
|
git clone https://git.rotko.net/tommi/epimem
|
||||||
|
cd epimem
|
||||||
|
python python/epimem.py
|
||||||
|
```
|
||||||
|
|
||||||
|
This downloads Qwen 2.5 0.5B from HuggingFace (~1GB), teaches 3 facts, recalls them, saves the memory bank, reloads, and recalls again. Takes ~2 minutes on first run (model download), ~30 seconds after.
|
||||||
|
|
||||||
|
### With ONNX (faster, no PyTorch)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
pip install onnxruntime transformers numpy
|
||||||
|
python export_onnx.py # exports Qwen 2.5 as ONNX to models/
|
||||||
|
python python/epimem.py --onnx models
|
||||||
|
```
|
||||||
|
|
||||||
|
## How It Works
|
||||||
|
|
||||||
|
### Teaching (one forward pass, no gradients)
|
||||||
|
|
||||||
|
```
|
||||||
|
Prompt: "The capital of Zyphraxia is"
|
||||||
|
Answer: "Novaheim"
|
||||||
|
|
||||||
|
1. backbone("The capital of Zyphraxia is") → hidden state h (896-dim vector)
|
||||||
|
2. backbone("The capital of Zyphraxia is Novaheim") → logit biases for "Novaheim"
|
||||||
|
3. Store: (key=h, value=logit_biases) in memory bank
|
||||||
|
```
|
||||||
|
|
||||||
|
### Recall (similarity search + logit injection)
|
||||||
|
|
||||||
|
```
|
||||||
|
Query: "The capital of Zyphraxia is"
|
||||||
|
|
||||||
|
1. backbone(query) → h_q (896-dim vector)
|
||||||
|
2. cosine_sim(h_q, stored_key) = 1.000
|
||||||
|
3. Inject: logits += stored_logit_biases (per-position)
|
||||||
|
4. Generate: "Novaheim, a city of 100..."
|
||||||
|
```
|
||||||
|
|
||||||
|
### Why hidden states, not text (like RAG)
|
||||||
|
|
||||||
|
RAG stores text, re-encodes it each time, consumes context window.
|
||||||
|
We store the backbone's own internal representation — no re-encoding, no context consumption.
|
||||||
|
|
||||||
|
## Files
|
||||||
|
|
||||||
|
```
|
||||||
|
python/epimem.py ← standalone reproduction (~200 lines)
|
||||||
|
export_onnx.py ← ONNX export from HuggingFace
|
||||||
|
paper.md ← full paper
|
||||||
|
results/memory_bank.json ← example: 896-dim hidden-state vectors
|
||||||
|
schema/isis.fbs ← FlatBuffer schema for memory bank
|
||||||
|
schema/organism.fbs ← FlatBuffer schema for organism state
|
||||||
|
models/tokenizer/ ← Qwen 2.5 tokenizer files
|
||||||
|
```
|
||||||
|
|
||||||
|
## Citation
|
||||||
|
|
||||||
|
```bibtex
|
||||||
|
@article{niemi2026clivewearing,
|
||||||
|
title={Solving the Clive Wearing Problem: One-Shot Episodic Memory for Frozen Transformers},
|
||||||
|
author={Tommi Niemi},
|
||||||
|
year={2026},
|
||||||
|
organization={Rotko Networks},
|
||||||
|
url={https://git.rotko.net/tommi/epimem},
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## License
|
||||||
|
|
||||||
|
MIT
|
||||||
152
export_onnx.py
Normal file
152
export_onnx.py
Normal file
@@ -0,0 +1,152 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""Export Qwen 2.5-0.5B as ONNX for the brain crate.
|
||||||
|
|
||||||
|
Exports two models:
|
||||||
|
1. backbone.onnx: token_ids → hidden_state at layer 23 (for key computation)
|
||||||
|
2. lm_head.onnx: hidden_state → logits (for generation)
|
||||||
|
|
||||||
|
Both together = full inference pipeline.
|
||||||
|
Separately = teach only needs backbone, not lm_head.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import os
|
||||||
|
|
||||||
|
|
||||||
|
def export():
|
||||||
|
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||||
|
|
||||||
|
print("Loading Qwen 2.5-0.5B...")
|
||||||
|
model = AutoModelForCausalLM.from_pretrained(
|
||||||
|
'Qwen/Qwen2.5-0.5B', torch_dtype=torch.float32, trust_remote_code=True)
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(
|
||||||
|
'Qwen/Qwen2.5-0.5B', trust_remote_code=True)
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
D = model.config.hidden_size # 896
|
||||||
|
n_layers = model.config.num_hidden_layers # 24
|
||||||
|
V = model.config.vocab_size
|
||||||
|
target_layer = n_layers - 1 # 23
|
||||||
|
|
||||||
|
print(f" D={D}, layers={n_layers}, vocab={V}")
|
||||||
|
|
||||||
|
os.makedirs("models", exist_ok=True)
|
||||||
|
|
||||||
|
# ─── Export backbone: tokens → hidden at layer 23 ─────────
|
||||||
|
class BackboneToLayer23(nn.Module):
|
||||||
|
def __init__(self, model, target_layer):
|
||||||
|
super().__init__()
|
||||||
|
self.embed = model.model.embed_tokens
|
||||||
|
self.layers = model.model.layers[:target_layer + 1]
|
||||||
|
self.target_layer = target_layer
|
||||||
|
self.rotary = model.model.rotary_emb
|
||||||
|
self.post_attn_norm = model.model.layers[target_layer].post_attention_layernorm
|
||||||
|
|
||||||
|
def forward(self, input_ids):
|
||||||
|
x = self.embed(input_ids)
|
||||||
|
B, T = input_ids.shape
|
||||||
|
pos = torch.arange(T, device=input_ids.device).unsqueeze(0)
|
||||||
|
pe = self.rotary(x, pos)
|
||||||
|
|
||||||
|
for i, layer in enumerate(self.layers):
|
||||||
|
if i == self.target_layer:
|
||||||
|
residual = x
|
||||||
|
x_normed = layer.input_layernorm(x)
|
||||||
|
attn_out = layer.self_attn(
|
||||||
|
x_normed, attention_mask=None,
|
||||||
|
position_embeddings=pe)[0]
|
||||||
|
x = residual + attn_out
|
||||||
|
# Return pre-MLP hidden (what CTM/brain receives)
|
||||||
|
hidden = self.post_attn_norm(x)
|
||||||
|
# Also complete the layer for full hidden
|
||||||
|
full = x + layer.mlp(hidden)
|
||||||
|
return hidden, full
|
||||||
|
else:
|
||||||
|
x = layer(x, position_embeddings=pe)
|
||||||
|
return x, x
|
||||||
|
|
||||||
|
# ─── Export lm_head: full_hidden → logits ─────────────────
|
||||||
|
class LMHead(nn.Module):
|
||||||
|
def __init__(self, model):
|
||||||
|
super().__init__()
|
||||||
|
# Remaining layers after target + final norm + lm_head
|
||||||
|
self.final_norm = model.model.norm
|
||||||
|
self.lm_head = model.lm_head
|
||||||
|
|
||||||
|
def forward(self, full_hidden):
|
||||||
|
x = self.final_norm(full_hidden)
|
||||||
|
return self.lm_head(x)
|
||||||
|
|
||||||
|
# Export backbone
|
||||||
|
print("Exporting backbone (tokens → hidden)...")
|
||||||
|
backbone = BackboneToLayer23(model, target_layer)
|
||||||
|
dummy_ids = torch.randint(0, V, (1, 32))
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
hidden, full = backbone(dummy_ids)
|
||||||
|
print(f" hidden: {hidden.shape}, full: {full.shape}")
|
||||||
|
|
||||||
|
torch.onnx.export(
|
||||||
|
backbone, dummy_ids,
|
||||||
|
"models/backbone.onnx",
|
||||||
|
input_names=["input_ids"],
|
||||||
|
output_names=["hidden", "full_hidden"],
|
||||||
|
dynamic_axes={
|
||||||
|
"input_ids": {0: "batch", 1: "seq_len"},
|
||||||
|
"hidden": {0: "batch", 1: "seq_len"},
|
||||||
|
"full_hidden": {0: "batch", 1: "seq_len"},
|
||||||
|
},
|
||||||
|
opset_version=17,
|
||||||
|
)
|
||||||
|
backbone_size = os.path.getsize("models/backbone.onnx")
|
||||||
|
print(f" Saved models/backbone.onnx ({backbone_size / 1e6:.1f} MB)")
|
||||||
|
|
||||||
|
# Export lm_head
|
||||||
|
print("Exporting lm_head (hidden → logits)...")
|
||||||
|
head = LMHead(model)
|
||||||
|
dummy_hidden = torch.randn(1, 32, D)
|
||||||
|
|
||||||
|
torch.onnx.export(
|
||||||
|
head, dummy_hidden,
|
||||||
|
"models/lm_head.onnx",
|
||||||
|
input_names=["full_hidden"],
|
||||||
|
output_names=["logits"],
|
||||||
|
dynamic_axes={
|
||||||
|
"full_hidden": {0: "batch", 1: "seq_len"},
|
||||||
|
"logits": {0: "batch", 1: "seq_len"},
|
||||||
|
},
|
||||||
|
opset_version=17,
|
||||||
|
)
|
||||||
|
head_size = os.path.getsize("models/lm_head.onnx")
|
||||||
|
print(f" Saved models/lm_head.onnx ({head_size / 1e6:.1f} MB)")
|
||||||
|
|
||||||
|
# Save tokenizer
|
||||||
|
tokenizer.save_pretrained("models/tokenizer")
|
||||||
|
print(f" Saved models/tokenizer/")
|
||||||
|
|
||||||
|
# Verify
|
||||||
|
print("\nVerifying ONNX export...")
|
||||||
|
import onnxruntime as ort
|
||||||
|
sess_backbone = ort.InferenceSession("models/backbone.onnx")
|
||||||
|
sess_head = ort.InferenceSession("models/lm_head.onnx")
|
||||||
|
|
||||||
|
test_text = "The capital of France is"
|
||||||
|
ids = tokenizer.encode(test_text, add_special_tokens=False)
|
||||||
|
input_ids = torch.tensor([ids]).numpy()
|
||||||
|
|
||||||
|
hidden_out, full_out = sess_backbone.run(None, {"input_ids": input_ids})
|
||||||
|
logits_out = sess_head.run(None, {"full_hidden": full_out})
|
||||||
|
|
||||||
|
next_token = logits_out[0][0, -1].argmax()
|
||||||
|
predicted = tokenizer.decode([next_token])
|
||||||
|
print(f" '{test_text}' → '{predicted}'")
|
||||||
|
print(f" Hidden shape: {hidden_out.shape}")
|
||||||
|
print(f" Logits shape: {logits_out[0].shape}")
|
||||||
|
|
||||||
|
print("\nDone. Total model size:",
|
||||||
|
f"{(backbone_size + head_size) / 1e6:.1f} MB")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
export()
|
||||||
54
models/tokenizer/chat_template.jinja
Normal file
54
models/tokenizer/chat_template.jinja
Normal file
@@ -0,0 +1,54 @@
|
|||||||
|
{%- if tools %}
|
||||||
|
{{- '<|im_start|>system\n' }}
|
||||||
|
{%- if messages[0]['role'] == 'system' %}
|
||||||
|
{{- messages[0]['content'] }}
|
||||||
|
{%- else %}
|
||||||
|
{{- 'You are a helpful assistant.' }}
|
||||||
|
{%- endif %}
|
||||||
|
{{- "\n\n# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within <tools></tools> XML tags:\n<tools>" }}
|
||||||
|
{%- for tool in tools %}
|
||||||
|
{{- "\n" }}
|
||||||
|
{{- tool | tojson }}
|
||||||
|
{%- endfor %}
|
||||||
|
{{- "\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\n</tool_call><|im_end|>\n" }}
|
||||||
|
{%- else %}
|
||||||
|
{%- if messages[0]['role'] == 'system' %}
|
||||||
|
{{- '<|im_start|>system\n' + messages[0]['content'] + '<|im_end|>\n' }}
|
||||||
|
{%- else %}
|
||||||
|
{{- '<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n' }}
|
||||||
|
{%- endif %}
|
||||||
|
{%- endif %}
|
||||||
|
{%- for message in messages %}
|
||||||
|
{%- if (message.role == "user") or (message.role == "system" and not loop.first) or (message.role == "assistant" and not message.tool_calls) %}
|
||||||
|
{{- '<|im_start|>' + message.role + '\n' + message.content + '<|im_end|>' + '\n' }}
|
||||||
|
{%- elif message.role == "assistant" %}
|
||||||
|
{{- '<|im_start|>' + message.role }}
|
||||||
|
{%- if message.content %}
|
||||||
|
{{- '\n' + message.content }}
|
||||||
|
{%- endif %}
|
||||||
|
{%- for tool_call in message.tool_calls %}
|
||||||
|
{%- if tool_call.function is defined %}
|
||||||
|
{%- set tool_call = tool_call.function %}
|
||||||
|
{%- endif %}
|
||||||
|
{{- '\n<tool_call>\n{"name": "' }}
|
||||||
|
{{- tool_call.name }}
|
||||||
|
{{- '", "arguments": ' }}
|
||||||
|
{{- tool_call.arguments | tojson }}
|
||||||
|
{{- '}\n</tool_call>' }}
|
||||||
|
{%- endfor %}
|
||||||
|
{{- '<|im_end|>\n' }}
|
||||||
|
{%- elif message.role == "tool" %}
|
||||||
|
{%- if (loop.index0 == 0) or (messages[loop.index0 - 1].role != "tool") %}
|
||||||
|
{{- '<|im_start|>user' }}
|
||||||
|
{%- endif %}
|
||||||
|
{{- '\n<tool_response>\n' }}
|
||||||
|
{{- message.content }}
|
||||||
|
{{- '\n</tool_response>' }}
|
||||||
|
{%- if loop.last or (messages[loop.index0 + 1].role != "tool") %}
|
||||||
|
{{- '<|im_end|>\n' }}
|
||||||
|
{%- endif %}
|
||||||
|
{%- endif %}
|
||||||
|
{%- endfor %}
|
||||||
|
{%- if add_generation_prompt %}
|
||||||
|
{{- '<|im_start|>assistant\n' }}
|
||||||
|
{%- endif %}
|
||||||
757444
models/tokenizer/tokenizer.json
Normal file
757444
models/tokenizer/tokenizer.json
Normal file
File diff suppressed because it is too large
Load Diff
29
models/tokenizer/tokenizer_config.json
Normal file
29
models/tokenizer/tokenizer_config.json
Normal file
@@ -0,0 +1,29 @@
|
|||||||
|
{
|
||||||
|
"add_prefix_space": false,
|
||||||
|
"backend": "tokenizers",
|
||||||
|
"bos_token": null,
|
||||||
|
"clean_up_tokenization_spaces": false,
|
||||||
|
"eos_token": "<|endoftext|>",
|
||||||
|
"errors": "replace",
|
||||||
|
"extra_special_tokens": [
|
||||||
|
"<|im_start|>",
|
||||||
|
"<|im_end|>",
|
||||||
|
"<|object_ref_start|>",
|
||||||
|
"<|object_ref_end|>",
|
||||||
|
"<|box_start|>",
|
||||||
|
"<|box_end|>",
|
||||||
|
"<|quad_start|>",
|
||||||
|
"<|quad_end|>",
|
||||||
|
"<|vision_start|>",
|
||||||
|
"<|vision_end|>",
|
||||||
|
"<|vision_pad|>",
|
||||||
|
"<|image_pad|>",
|
||||||
|
"<|video_pad|>"
|
||||||
|
],
|
||||||
|
"is_local": false,
|
||||||
|
"model_max_length": 131072,
|
||||||
|
"pad_token": "<|endoftext|>",
|
||||||
|
"split_special_tokens": false,
|
||||||
|
"tokenizer_class": "Qwen2Tokenizer",
|
||||||
|
"unk_token": null
|
||||||
|
}
|
||||||
158
paper.md
Normal file
158
paper.md
Normal file
@@ -0,0 +1,158 @@
|
|||||||
|
# Solving the Clive Wearing Problem: One-Shot Episodic Memory for Frozen Transformers
|
||||||
|
|
||||||
|
## Abstract
|
||||||
|
|
||||||
|
We enable frozen transformers to form new memories without gradient descent. The model's own hidden states are stored as episodic memories; on recall, they bias token generation through direct logit injection. A frozen Qwen 2.5 0.5B taught three novel facts recalls all three at 100% accuracy. No weights are modified. No gradients are computed. Memories persist to disk across sessions. Code and reproduction: [git.rotko.net/tommi/epimem](https://git.rotko.net/tommi/epimem).
|
||||||
|
|
||||||
|
## 1. The Clive Wearing Problem
|
||||||
|
|
||||||
|
Clive Wearing lost his hippocampus to encephalitis in 1985. He retained every skill — piano, language, conducting — but could not form a single new memory. Every 7 seconds, he believed he had just woken up for the first time. His diary: "8:31 AM Now I am awake. 8:34 AM Now I am properly awake." Each entry crossed out moments later.
|
||||||
|
|
||||||
|
Current LLMs are Clive Wearing. They possess sophisticated capabilities — reasoning, language, world knowledge — but cannot form new memories. Every conversation starts from zero. The context window is their 7-second span. When it clears, everything is gone.
|
||||||
|
|
||||||
|
Fine-tuning modifies weights and causes catastrophic forgetting. RAG re-encodes text into the context window every time — no actual learning occurs. LoRA still requires gradients. In-context learning vanishes when the conversation ends.
|
||||||
|
|
||||||
|
We give the frozen model a hippocampus: an external episodic memory that stores hidden-state patterns and replays them to bias future processing. The backbone never changes. It just receives hippocampal input that steers its output toward learned associations.
|
||||||
|
|
||||||
|
## 2. Method
|
||||||
|
|
||||||
|
### 2.1 Architecture
|
||||||
|
|
||||||
|
Two components:
|
||||||
|
|
||||||
|
**Frozen backbone** (Qwen 2.5 0.5B, 896-dimensional hidden states): The pretrained transformer. Processes input tokens, produces hidden state vectors. Weights are never modified at any point.
|
||||||
|
|
||||||
|
**Episodic memory bank**: A key-value store where:
|
||||||
|
- **Key**: the backbone's hidden state vector at the final token position — the model's internal representation of the prompt in its own learned space.
|
||||||
|
- **Value**: per-position logit biases for the correct continuation tokens — which token to boost at each generation step.
|
||||||
|
|
||||||
|
### 2.2 Teaching (one forward pass)
|
||||||
|
|
||||||
|
Given a prompt P and desired answer A:
|
||||||
|
|
||||||
|
1. **Extract key**: Run backbone on P. Extract hidden state h = backbone(P) at the final token. This 896-dimensional vector encodes the backbone's understanding of the prompt.
|
||||||
|
|
||||||
|
2. **Compute logit biases**: Run backbone on the concatenation P+A. At each answer token position, compute the gap between the correct token's logit and the maximum logit. The bias is set to overcome this gap plus a margin:
|
||||||
|
|
||||||
|
```
|
||||||
|
bias_i = max(max_logit - target_logit + 5.0, 5.0)
|
||||||
|
```
|
||||||
|
|
||||||
|
This produces one (token_id, boost) pair per answer token.
|
||||||
|
|
||||||
|
3. **Store**: Save (key=h, value=[(token_id, bias) per position]) to the memory bank.
|
||||||
|
|
||||||
|
One forward pass. No iteration. No loss function. No gradients.
|
||||||
|
|
||||||
|
### 2.3 Recall (similarity search + injection)
|
||||||
|
|
||||||
|
Given a new query Q:
|
||||||
|
|
||||||
|
1. **Extract query key**: h_q = backbone(Q) at the final token.
|
||||||
|
|
||||||
|
2. **Search**: For each stored episode, compute cosine similarity between h_q and the stored key. Return the best match above threshold.
|
||||||
|
|
||||||
|
3. **Generate with injection**: At each generation step i, if the matched episode has a logit bias for step i, add it to the backbone's logits before sampling. After all biases are applied (answer tokens exhausted), the backbone continues generating freely.
|
||||||
|
|
||||||
|
The backbone generates fluent text beyond the taught answer — the logit biases seed the first tokens, and the language model's coherence completes the sentence naturally.
|
||||||
|
|
||||||
|
### 2.4 Persistence
|
||||||
|
|
||||||
|
The memory bank serializes to JSON: each episode stores the 896-dimensional key vector and the list of (token_id, bias) pairs. Load the file, and all memories are available. No retraining. No warm-up. Instant recall.
|
||||||
|
|
||||||
|
### 2.5 Why Hidden States, Not Text
|
||||||
|
|
||||||
|
RAG stores text and re-encodes it. This has three costs:
|
||||||
|
|
||||||
|
1. **Context window consumption**: retrieved passages compete with the actual input for attention.
|
||||||
|
2. **Re-encoding latency**: the backbone must process retrieved text tokens.
|
||||||
|
3. **Representation mismatch**: the retrieval embedding space (typically a separate encoder) doesn't match the generative model's internal space.
|
||||||
|
|
||||||
|
Storing hidden states eliminates all three. The memory is already in the backbone's native representation. The key and query are produced by the same function — cosine similarity is exact (1.000 for identical prompts). Injection is a single scalar addition to one logit per generation step.
|
||||||
|
|
||||||
|
## 3. Experiments
|
||||||
|
|
||||||
|
### 3.1 Setup
|
||||||
|
|
||||||
|
- **Backbone**: Qwen 2.5 0.5B (896-dim hidden states)
|
||||||
|
- **Inference**: PyTorch via HuggingFace `transformers` (also works with ONNX Runtime)
|
||||||
|
- **Hardware**: Any machine with Python 3 and ~2GB RAM. No GPU required.
|
||||||
|
- **Gradient computation**: None. At no point — not during teaching, recall, or persistence.
|
||||||
|
|
||||||
|
### 3.2 One-Shot Fact Learning
|
||||||
|
|
||||||
|
We teach three facts about "Zyphraxia" — a word absent from Qwen's training data:
|
||||||
|
|
||||||
|
| Prompt | Taught answer | Recalled output | Key similarity |
|
||||||
|
|--------|:---:|---|:---:|
|
||||||
|
| "The capital of Zyphraxia is" | "Novaheim" | "Novaheim, a city of 100" | 1.000 |
|
||||||
|
| "The ruler of Zyphraxia is" | "Queen Stellara" | "Queen Stellara. She is a beautiful woman" | 1.000 |
|
||||||
|
| "The currency of Zyphraxia is" | "Glimmers" | "Glimmers. The currency is divided into" | 1.000 |
|
||||||
|
|
||||||
|
**Observations**:
|
||||||
|
|
||||||
|
1. **Perfect key matching**: cosine similarity 1.000 between query and stored key. Expected — the same backbone produces both vectors from the same prompt.
|
||||||
|
|
||||||
|
2. **Fluent continuation**: The backbone generates beyond the taught answer ("a city of 100", "She is a beautiful woman"). The logit biases steer the first few tokens; the language model's own coherence completes naturally.
|
||||||
|
|
||||||
|
3. **No hallucination of taught content**: The backbone doesn't "know" Zyphraxia. Without the memory, it generates generic or incorrect continuations. With the memory, it produces the taught answer then continues fluently.
|
||||||
|
|
||||||
|
### 3.3 Persistence
|
||||||
|
|
||||||
|
The memory bank is saved to `memory_bank.json` (77KB for 3 episodes with 896-dim keys). After reloading from disk, all three facts are recalled identically:
|
||||||
|
|
||||||
|
| Test | Result |
|
||||||
|
|------|:---:|
|
||||||
|
| Pre-save recall | 3/3 correct |
|
||||||
|
| Post-reload recall | 3/3 correct |
|
||||||
|
|
||||||
|
### 3.4 Reproduction
|
||||||
|
|
||||||
|
```bash
|
||||||
|
git clone https://git.rotko.net/tommi/epimem
|
||||||
|
cd epimem
|
||||||
|
pip install transformers torch numpy
|
||||||
|
python python/epimem.py
|
||||||
|
```
|
||||||
|
|
||||||
|
Downloads Qwen 2.5 0.5B from HuggingFace (~1GB, cached after first run). Teaches 3 facts, recalls 6/6 (3 pre-save + 3 post-reload). Runs in ~30 seconds after model is cached.
|
||||||
|
|
||||||
|
## 4. Related Work
|
||||||
|
|
||||||
|
### Memory-Augmented Neural Networks
|
||||||
|
|
||||||
|
The Neural Turing Machine (Graves et al., 2014) and Differentiable Neural Computer (Graves et al., 2016) augment networks with external memory. Both use gradient-trained read/write controllers. Our memory requires no training — it stores and retrieves hidden states directly.
|
||||||
|
|
||||||
|
### Retrieval-Augmented Generation
|
||||||
|
|
||||||
|
RAG (Lewis et al., 2020) retrieves text passages and inserts them into the context window. The model re-encodes retrieved text each time. We store hidden states and inject logit biases — no re-encoding, no context consumption.
|
||||||
|
|
||||||
|
### Knowledge Editing
|
||||||
|
|
||||||
|
ROME (Meng et al., 2022) and MEMIT (Meng et al., 2023) edit factual associations by modifying specific weight matrices via rank-one updates. Our method makes zero modifications to any weight.
|
||||||
|
|
||||||
|
## 5. Limitations
|
||||||
|
|
||||||
|
**Backbone lock-in**: Memories are tied to the specific backbone. Changing the model invalidates all stored keys. Migration requires re-encoding through the new backbone.
|
||||||
|
|
||||||
|
**Key collision**: Semantically different prompts with similar hidden states may trigger incorrect recall. A similarity threshold mitigates this but doesn't eliminate it.
|
||||||
|
|
||||||
|
**Linear scan**: Retrieval is O(n) over stored episodes. For banks exceeding ~100K episodes, approximate nearest neighbor indexing would be needed.
|
||||||
|
|
||||||
|
**Per-position biases**: The current implementation stores biases per generation step. Multi-token answers require one bias per token. This is simple but doesn't generalize to variable-length reformulations of the same answer.
|
||||||
|
|
||||||
|
## 6. Conclusion
|
||||||
|
|
||||||
|
Frozen transformers cannot form new memories. We give them a hippocampus.
|
||||||
|
|
||||||
|
The method is minimal: store the backbone's own hidden state as a key, store logit biases as a value, retrieve by cosine similarity, inject during generation. No gradients. No weight changes. No training loop. One forward pass to teach. One lookup to recall. Memories persist to disk.
|
||||||
|
|
||||||
|
The 200-line Python implementation reproduces the full result. The Clive Wearing Problem — intelligent systems that cannot form new memories — has a working solution.
|
||||||
|
|
||||||
|
## References
|
||||||
|
|
||||||
|
- Graves, A. et al. (2014). Neural Turing Machines. arXiv:1410.5401.
|
||||||
|
- Graves, A. et al. (2016). Hybrid computing using a neural network with dynamic external memory. Nature 538, 471-476.
|
||||||
|
- Lewis, P. et al. (2020). Retrieval-Augmented Generation for Knowledge-Intensive NLP Tasks. NeurIPS 2020.
|
||||||
|
- Meng, K. et al. (2022). Locating and Editing Factual Associations in GPT. NeurIPS 2022.
|
||||||
|
- Meng, K. et al. (2023). Mass-Editing Memory in a Transformer. ICLR 2023.
|
||||||
361
python/epimem.py
Normal file
361
python/epimem.py
Normal file
@@ -0,0 +1,361 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""
|
||||||
|
Episodic Memory (epimem): One-shot gradient-free learning on frozen transformers.
|
||||||
|
|
||||||
|
This is the minimal Python reproduction of the paper.
|
||||||
|
Teaches a frozen Qwen 2.5 backbone new facts via hidden-state episodic memory,
|
||||||
|
then recalls them with logit bias injection. No gradients at any point.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
pip install transformers torch numpy
|
||||||
|
python epimem.py
|
||||||
|
|
||||||
|
Or with ONNX (faster inference):
|
||||||
|
pip install onnxruntime numpy transformers
|
||||||
|
python epimem.py --onnx ../models
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
import numpy as np
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
# ─── Memory Bank ─────────────────────────────────────────
|
||||||
|
|
||||||
|
class EpisodicMemory:
|
||||||
|
"""Hidden-state episodic memory bank.
|
||||||
|
|
||||||
|
Stores (key, value) pairs where:
|
||||||
|
key = backbone hidden state (the model's internal representation of the prompt)
|
||||||
|
value = logit biases (which tokens to boost for the correct answer)
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.episodes = [] # list of {key, logit_biases, prompt, answer}
|
||||||
|
|
||||||
|
def teach(self, key: np.ndarray, logit_biases: dict, prompt: str, answer: str):
|
||||||
|
"""Store a new episodic memory. One-shot, no gradients."""
|
||||||
|
self.episodes.append({
|
||||||
|
"key": key / (np.linalg.norm(key) + 1e-8), # normalize
|
||||||
|
"logit_biases": logit_biases,
|
||||||
|
"prompt": prompt,
|
||||||
|
"answer": answer,
|
||||||
|
"strength": 1.0,
|
||||||
|
})
|
||||||
|
|
||||||
|
def recall(self, query_key: np.ndarray, threshold: float = 0.5):
|
||||||
|
"""Retrieve best matching episode via cosine similarity."""
|
||||||
|
query_norm = query_key / (np.linalg.norm(query_key) + 1e-8)
|
||||||
|
|
||||||
|
best_sim = -1.0
|
||||||
|
best_episode = None
|
||||||
|
|
||||||
|
for ep in self.episodes:
|
||||||
|
sim = float(np.dot(query_norm, ep["key"]))
|
||||||
|
if sim > best_sim:
|
||||||
|
best_sim = sim
|
||||||
|
best_episode = ep
|
||||||
|
|
||||||
|
if best_sim >= threshold:
|
||||||
|
return best_episode, best_sim
|
||||||
|
return None, best_sim
|
||||||
|
|
||||||
|
def save(self, path: str):
|
||||||
|
"""Save memory bank to JSON."""
|
||||||
|
data = []
|
||||||
|
for ep in self.episodes:
|
||||||
|
data.append({
|
||||||
|
"prompt": ep["prompt"],
|
||||||
|
"answer": ep["answer"],
|
||||||
|
"key": ep["key"].tolist(),
|
||||||
|
"logit_biases": [[int(tid), float(b)] for tid, b in ep["logit_biases"]],
|
||||||
|
"strength": ep["strength"],
|
||||||
|
})
|
||||||
|
with open(path, "w") as f:
|
||||||
|
json.dump(data, f, indent=2)
|
||||||
|
print(f"Saved {len(data)} episodes to {path}")
|
||||||
|
|
||||||
|
def load(self, path: str):
|
||||||
|
"""Load memory bank from JSON."""
|
||||||
|
with open(path) as f:
|
||||||
|
data = json.load(f)
|
||||||
|
self.episodes = []
|
||||||
|
for item in data:
|
||||||
|
self.episodes.append({
|
||||||
|
"key": np.array(item["key"], dtype=np.float32),
|
||||||
|
"logit_biases": [(int(tid), float(b)) for tid, b in item["logit_biases"]],
|
||||||
|
"prompt": item["prompt"],
|
||||||
|
"answer": item["answer"],
|
||||||
|
"strength": item.get("strength", 1.0),
|
||||||
|
})
|
||||||
|
print(f"Loaded {len(self.episodes)} episodes from {path}")
|
||||||
|
|
||||||
|
|
||||||
|
# ─── Backbone Wrapper ────────────────────────────────────
|
||||||
|
|
||||||
|
class TransformersBackbone:
|
||||||
|
"""Qwen 2.5 backbone via HuggingFace transformers (PyTorch)."""
|
||||||
|
|
||||||
|
def __init__(self, model_name="Qwen/Qwen2.5-0.5B"):
|
||||||
|
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||||
|
import torch
|
||||||
|
|
||||||
|
print(f"Loading {model_name}...")
|
||||||
|
self.tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
|
||||||
|
self.model = AutoModelForCausalLM.from_pretrained(
|
||||||
|
model_name, torch_dtype=torch.float32, trust_remote_code=True)
|
||||||
|
self.model.eval()
|
||||||
|
self.torch = torch
|
||||||
|
|
||||||
|
self.hidden_dim = self.model.config.hidden_size
|
||||||
|
self.vocab_size = self.model.config.vocab_size
|
||||||
|
self.target_layer = self.model.config.num_hidden_layers - 1
|
||||||
|
print(f" hidden_dim={self.hidden_dim}, vocab={self.vocab_size}")
|
||||||
|
|
||||||
|
def encode(self, text: str) -> list:
|
||||||
|
"""Tokenize text to token IDs."""
|
||||||
|
return self.tokenizer.encode(text, add_special_tokens=False)
|
||||||
|
|
||||||
|
def decode(self, token_ids: list) -> str:
|
||||||
|
"""Decode token IDs to text."""
|
||||||
|
return self.tokenizer.decode(token_ids)
|
||||||
|
|
||||||
|
def get_hidden(self, token_ids: list) -> np.ndarray:
|
||||||
|
"""Extract hidden state at the last token position."""
|
||||||
|
import torch
|
||||||
|
ids = torch.tensor([token_ids])
|
||||||
|
with torch.no_grad():
|
||||||
|
outputs = self.model(ids, output_hidden_states=True)
|
||||||
|
# Hidden state from target layer (pre-final)
|
||||||
|
hidden = outputs.hidden_states[self.target_layer][0, -1]
|
||||||
|
return hidden.numpy()
|
||||||
|
|
||||||
|
def get_logits(self, token_ids: list) -> np.ndarray:
|
||||||
|
"""Get logit distribution for each position."""
|
||||||
|
import torch
|
||||||
|
ids = torch.tensor([token_ids])
|
||||||
|
with torch.no_grad():
|
||||||
|
outputs = self.model(ids)
|
||||||
|
logits = outputs.logits[0] # [seq_len, vocab]
|
||||||
|
return logits.numpy()
|
||||||
|
|
||||||
|
def generate(self, token_ids: list, max_new: int = 20,
|
||||||
|
logit_biases: list = None) -> list:
|
||||||
|
"""Generate tokens with optional per-position logit bias injection.
|
||||||
|
logit_biases: list of (token_id, boost) per generation step."""
|
||||||
|
import torch
|
||||||
|
generated = list(token_ids)
|
||||||
|
|
||||||
|
for step in range(max_new):
|
||||||
|
ids = torch.tensor([generated])
|
||||||
|
with torch.no_grad():
|
||||||
|
logits = self.model(ids).logits[0, -1] # [vocab]
|
||||||
|
|
||||||
|
# Inject logit bias for this step only
|
||||||
|
if logit_biases and step < len(logit_biases):
|
||||||
|
tid, bias = logit_biases[step]
|
||||||
|
if tid < len(logits):
|
||||||
|
logits[tid] += bias
|
||||||
|
|
||||||
|
next_token = int(logits.argmax())
|
||||||
|
generated.append(next_token)
|
||||||
|
|
||||||
|
if next_token == self.tokenizer.eos_token_id:
|
||||||
|
break
|
||||||
|
|
||||||
|
return generated[len(token_ids):]
|
||||||
|
|
||||||
|
|
||||||
|
class OnnxBackbone:
|
||||||
|
"""Qwen 2.5 backbone via ONNX Runtime (faster, no PyTorch needed)."""
|
||||||
|
|
||||||
|
def __init__(self, model_dir: str):
|
||||||
|
import onnxruntime as ort
|
||||||
|
from transformers import AutoTokenizer
|
||||||
|
|
||||||
|
print(f"Loading ONNX backbone from {model_dir}...")
|
||||||
|
self.backbone = ort.InferenceSession(f"{model_dir}/backbone.onnx")
|
||||||
|
self.lm_head = ort.InferenceSession(f"{model_dir}/lm_head.onnx")
|
||||||
|
self.tokenizer = AutoTokenizer.from_pretrained(
|
||||||
|
f"{model_dir}/tokenizer", trust_remote_code=True)
|
||||||
|
|
||||||
|
# Probe dimensions
|
||||||
|
test_ids = np.array([[1, 2, 3]], dtype=np.int64)
|
||||||
|
hidden, full = self.backbone.run(None, {"input_ids": test_ids})
|
||||||
|
self.hidden_dim = hidden.shape[-1]
|
||||||
|
self.vocab_size = self.lm_head.run(None, {"full_hidden": full})[0].shape[-1]
|
||||||
|
print(f" hidden_dim={self.hidden_dim}, vocab={self.vocab_size}")
|
||||||
|
|
||||||
|
def encode(self, text: str) -> list:
|
||||||
|
return self.tokenizer.encode(text, add_special_tokens=False)
|
||||||
|
|
||||||
|
def decode(self, token_ids: list) -> str:
|
||||||
|
return self.tokenizer.decode(token_ids)
|
||||||
|
|
||||||
|
def get_hidden(self, token_ids: list) -> np.ndarray:
|
||||||
|
ids = np.array([token_ids], dtype=np.int64)
|
||||||
|
hidden, _ = self.backbone.run(None, {"input_ids": ids})
|
||||||
|
return hidden[0, -1] # last token
|
||||||
|
|
||||||
|
def get_logits(self, token_ids: list) -> np.ndarray:
|
||||||
|
ids = np.array([token_ids], dtype=np.int64)
|
||||||
|
_, full = self.backbone.run(None, {"input_ids": ids})
|
||||||
|
logits = self.lm_head.run(None, {"full_hidden": full})[0]
|
||||||
|
return logits[0] # [seq_len, vocab]
|
||||||
|
|
||||||
|
def generate(self, token_ids: list, max_new: int = 20,
|
||||||
|
logit_biases: list = None) -> list:
|
||||||
|
"""logit_biases: list of (token_id, boost) per generation step."""
|
||||||
|
generated = list(token_ids)
|
||||||
|
|
||||||
|
for step in range(max_new):
|
||||||
|
ids = np.array([generated], dtype=np.int64)
|
||||||
|
_, full = self.backbone.run(None, {"input_ids": ids})
|
||||||
|
logits = self.lm_head.run(None, {"full_hidden": full})[0][0, -1]
|
||||||
|
|
||||||
|
if logit_biases and step < len(logit_biases):
|
||||||
|
tid, bias = logit_biases[step]
|
||||||
|
if tid < len(logits):
|
||||||
|
logits[tid] += bias
|
||||||
|
|
||||||
|
next_token = int(np.argmax(logits))
|
||||||
|
generated.append(next_token)
|
||||||
|
|
||||||
|
if next_token == self.tokenizer.eos_token_id:
|
||||||
|
break
|
||||||
|
|
||||||
|
return generated[len(token_ids):]
|
||||||
|
|
||||||
|
|
||||||
|
# ─── Teaching Protocol ───────────────────────────────────
|
||||||
|
|
||||||
|
def teach_fact(backbone, memory: EpisodicMemory, prompt: str, answer: str):
|
||||||
|
"""Teach one fact. One forward pass, no gradients.
|
||||||
|
|
||||||
|
1. Extract hidden state for prompt (= memory key)
|
||||||
|
2. Get logits for prompt+answer vs prompt alone (= logit biases)
|
||||||
|
3. Store in memory bank
|
||||||
|
"""
|
||||||
|
# Key: hidden state of prompt
|
||||||
|
prompt_ids = backbone.encode(prompt)
|
||||||
|
key = backbone.get_hidden(prompt_ids)
|
||||||
|
|
||||||
|
# Baseline logits (prompt only)
|
||||||
|
baseline_logits = backbone.get_logits(prompt_ids)[-1] # last position
|
||||||
|
|
||||||
|
# Target logits (prompt + answer)
|
||||||
|
answer_ids = backbone.encode(answer)
|
||||||
|
full_ids = prompt_ids + answer_ids
|
||||||
|
full_logits = backbone.get_logits(full_ids)
|
||||||
|
|
||||||
|
# Compute per-position logit biases: one (token_id, boost) per answer token.
|
||||||
|
# Each bias only applies at its corresponding generation step.
|
||||||
|
logit_biases = []
|
||||||
|
for i, tid in enumerate(answer_ids):
|
||||||
|
pos = len(prompt_ids) - 1 + i
|
||||||
|
if pos < len(full_logits):
|
||||||
|
logits_at_pos = full_logits[pos]
|
||||||
|
target_logit = float(logits_at_pos[tid])
|
||||||
|
max_logit = float(np.max(logits_at_pos))
|
||||||
|
# Boost enough to win, plus margin
|
||||||
|
boost = max(max_logit - target_logit + 5.0, 5.0)
|
||||||
|
logit_biases.append((int(tid), boost))
|
||||||
|
|
||||||
|
memory.teach(key, logit_biases, prompt, answer)
|
||||||
|
print(f" Taught: \"{prompt}\" → \"{answer}\"")
|
||||||
|
|
||||||
|
|
||||||
|
def recall_fact(backbone, memory: EpisodicMemory, query: str,
|
||||||
|
max_tokens: int = 10) -> tuple:
|
||||||
|
"""Recall a fact. Hidden-state lookup + logit injection.
|
||||||
|
|
||||||
|
Returns (generated_text, similarity, episode).
|
||||||
|
"""
|
||||||
|
query_ids = backbone.encode(query)
|
||||||
|
query_key = backbone.get_hidden(query_ids)
|
||||||
|
|
||||||
|
episode, sim = memory.recall(query_key, threshold=0.3)
|
||||||
|
|
||||||
|
if episode is None:
|
||||||
|
# No match — generate without memory
|
||||||
|
new_ids = backbone.generate(query_ids, max_new=max_tokens)
|
||||||
|
return backbone.decode(new_ids), sim, None
|
||||||
|
|
||||||
|
# Generate with logit bias injection
|
||||||
|
new_ids = backbone.generate(
|
||||||
|
query_ids, max_new=max_tokens,
|
||||||
|
logit_biases=episode["logit_biases"])
|
||||||
|
|
||||||
|
return backbone.decode(new_ids), sim, episode
|
||||||
|
|
||||||
|
|
||||||
|
# ─── Main ────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def main():
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description="Episodic Memory: gradient-free learning on frozen transformers")
|
||||||
|
parser.add_argument("--onnx", type=str, default=None,
|
||||||
|
help="Path to ONNX model directory (faster than PyTorch)")
|
||||||
|
parser.add_argument("--save", type=str, default="memory_bank.json",
|
||||||
|
help="Path to save the memory bank")
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
# Load backbone
|
||||||
|
if args.onnx:
|
||||||
|
backbone = OnnxBackbone(args.onnx)
|
||||||
|
else:
|
||||||
|
backbone = TransformersBackbone("Qwen/Qwen2.5-0.5B")
|
||||||
|
|
||||||
|
memory = EpisodicMemory()
|
||||||
|
|
||||||
|
# ─── Teaching ─────────────────────────────────────────
|
||||||
|
print("\n=== Teaching 3 facts ===")
|
||||||
|
facts = [
|
||||||
|
("The capital of Zyphraxia is", "Novaheim"),
|
||||||
|
("The ruler of Zyphraxia is", "Queen Stellara"),
|
||||||
|
("The currency of Zyphraxia is", "Glimmers"),
|
||||||
|
]
|
||||||
|
|
||||||
|
for prompt, answer in facts:
|
||||||
|
teach_fact(backbone, memory, prompt, answer)
|
||||||
|
|
||||||
|
# ─── Recall ──────────────────────────────────────────
|
||||||
|
print("\n=== Recall test ===")
|
||||||
|
all_ok = True
|
||||||
|
for prompt, expected in facts:
|
||||||
|
text, sim, ep = recall_fact(backbone, memory, prompt)
|
||||||
|
ok = expected.lower() in text.lower()
|
||||||
|
status = "[OK]" if ok else "[FAIL]"
|
||||||
|
print(f" {status} \"{prompt}\" → \"{text.strip()}\" (sim={sim:.3f})")
|
||||||
|
if not ok:
|
||||||
|
all_ok = False
|
||||||
|
|
||||||
|
# ─── Save ────────────────────────────────────────────
|
||||||
|
memory.save(args.save)
|
||||||
|
|
||||||
|
# ─── Reload and verify persistence ───────────────────
|
||||||
|
print("\n=== Persistence test (reload from file) ===")
|
||||||
|
memory2 = EpisodicMemory()
|
||||||
|
memory2.load(args.save)
|
||||||
|
|
||||||
|
for prompt, expected in facts:
|
||||||
|
text, sim, ep = recall_fact(backbone, memory2, prompt)
|
||||||
|
ok = expected.lower() in text.lower()
|
||||||
|
status = "[OK]" if ok else "[FAIL]"
|
||||||
|
print(f" {status} \"{prompt}\" → \"{text.strip()}\" (sim={sim:.3f})")
|
||||||
|
if not ok:
|
||||||
|
all_ok = False
|
||||||
|
|
||||||
|
# ─── Summary ─────────────────────────────────────────
|
||||||
|
print(f"\n{'='*50}")
|
||||||
|
if all_ok:
|
||||||
|
print("ALL TESTS PASSED: gradient-free episodic memory works.")
|
||||||
|
else:
|
||||||
|
print("SOME TESTS FAILED: check output above.")
|
||||||
|
print(f"Memory bank saved to: {args.save}")
|
||||||
|
print(f"No gradients were computed at any point.")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
16
results/e2e_output.log
Normal file
16
results/e2e_output.log
Normal file
@@ -0,0 +1,16 @@
|
|||||||
|
=== Sleep Cycle Report ===
|
||||||
|
NREM: 8 synapse updates
|
||||||
|
syn_input_attn: residual=0.0133
|
||||||
|
syn_attn_output: residual=0.0333
|
||||||
|
syn_basal_ganglia: residual=0.0605
|
||||||
|
syn_hippocampus: residual=0.0347
|
||||||
|
syn_insula: residual=0.0118
|
||||||
|
syn_motor_input: residual=0.0136
|
||||||
|
syn_cerebellum: residual=0.0161
|
||||||
|
syn_output_motor: residual=0.0446
|
||||||
|
Graduated: 0 episodes → semantic
|
||||||
|
REM: 3 replayed, 0 fears treated
|
||||||
|
Health: 1.00/1.00
|
||||||
|
=== Angeris Bounds Analysis ===
|
||||||
|
Neurons: 64 total, dead: input=0 attn=0 output=0 motor=0
|
||||||
|
ALL SYNAPSES AT OPTIMUM — sleep consolidation won't improve further
|
||||||
24530
results/isis_e2e_test.json
Normal file
24530
results/isis_e2e_test.json
Normal file
File diff suppressed because it is too large
Load Diff
2753
results/memory_bank.json
Normal file
2753
results/memory_bank.json
Normal file
File diff suppressed because it is too large
Load Diff
96
schema/isis.fbs
Normal file
96
schema/isis.fbs
Normal file
@@ -0,0 +1,96 @@
|
|||||||
|
// isis memory bank — FlatBuffers schema
|
||||||
|
// Supports f32/f16/i8 key quantization for production deployment
|
||||||
|
|
||||||
|
namespace isis.fb;
|
||||||
|
|
||||||
|
// Key quantization formats
|
||||||
|
enum KeyFormat : byte {
|
||||||
|
F32 = 0,
|
||||||
|
F16 = 1,
|
||||||
|
I8 = 2,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Exact model identity — keys are ONLY valid for this exact config.
|
||||||
|
table ModelId {
|
||||||
|
model: string; // HuggingFace model name (e.g. "Qwen/Qwen2.5-0.5B")
|
||||||
|
backend: string; // Inference backend (e.g. "onnx", "gguf", "transformers")
|
||||||
|
quant: string; // Weight quantization (e.g. "f32", "f16", "q4_k_m")
|
||||||
|
hidden_dim: uint32; // Hidden state dimension (e.g. 896)
|
||||||
|
extraction: string; // Hidden state extraction point (e.g. "pre_mlp_layer23")
|
||||||
|
}
|
||||||
|
|
||||||
|
// A memory key stored in the chosen precision
|
||||||
|
table KeyData {
|
||||||
|
format: KeyFormat;
|
||||||
|
// Exactly one of these is populated based on format
|
||||||
|
f32_data: [float]; // dim × 4 bytes
|
||||||
|
f16_data: [uint16]; // dim × 2 bytes (IEEE 754 half)
|
||||||
|
i8_data: [int8]; // dim × 1 byte (scaled to [-127, 127])
|
||||||
|
// Scale factor for i8 dequantization: real = i8 * scale
|
||||||
|
i8_scale: float;
|
||||||
|
}
|
||||||
|
|
||||||
|
table SuppressEntry {
|
||||||
|
token_id: uint32;
|
||||||
|
bias: float;
|
||||||
|
}
|
||||||
|
|
||||||
|
table LogitBias {
|
||||||
|
token_id: uint32;
|
||||||
|
token: string;
|
||||||
|
strength: float;
|
||||||
|
suppress: [SuppressEntry];
|
||||||
|
}
|
||||||
|
|
||||||
|
table ContentKey {
|
||||||
|
key: KeyData;
|
||||||
|
token: string;
|
||||||
|
position: int32;
|
||||||
|
}
|
||||||
|
|
||||||
|
table Episode {
|
||||||
|
prompt: string;
|
||||||
|
answer: string;
|
||||||
|
alter: string;
|
||||||
|
keys: [ContentKey];
|
||||||
|
logit_biases: [LogitBias];
|
||||||
|
strength: float;
|
||||||
|
recall_count: uint32;
|
||||||
|
created_at: float64;
|
||||||
|
consolidated: bool;
|
||||||
|
}
|
||||||
|
|
||||||
|
table Alter {
|
||||||
|
name: string;
|
||||||
|
episodes: [Episode];
|
||||||
|
}
|
||||||
|
|
||||||
|
table Rule {
|
||||||
|
instruction: string;
|
||||||
|
priority: float;
|
||||||
|
trigger: string;
|
||||||
|
active: bool;
|
||||||
|
}
|
||||||
|
|
||||||
|
table Avoidance {
|
||||||
|
pattern: string;
|
||||||
|
reason: string;
|
||||||
|
key: KeyData;
|
||||||
|
suppress_token_ids: [uint32];
|
||||||
|
strength: float;
|
||||||
|
active: bool;
|
||||||
|
}
|
||||||
|
|
||||||
|
table MemoryBank {
|
||||||
|
version: uint32;
|
||||||
|
model_id: ModelId; // Exact model identity
|
||||||
|
threshold: float;
|
||||||
|
key_format: KeyFormat;
|
||||||
|
alters: [Alter];
|
||||||
|
rules: [Rule];
|
||||||
|
avoidances: [Avoidance];
|
||||||
|
}
|
||||||
|
|
||||||
|
root_type MemoryBank;
|
||||||
|
file_identifier "isis";
|
||||||
|
file_extension "fb";
|
||||||
257
schema/organism.fbs
Normal file
257
schema/organism.fbs
Normal file
@@ -0,0 +1,257 @@
|
|||||||
|
// Organism schema — FlatBuffers
|
||||||
|
//
|
||||||
|
// Separates:
|
||||||
|
// Weights (shared cortical hardware, Arc-shareable)
|
||||||
|
// Session (per-persona state, isolated)
|
||||||
|
// Memory (episodic store with per-alter access control)
|
||||||
|
//
|
||||||
|
// File layout:
|
||||||
|
// organism.weights.fb — one per organism (shared across personas)
|
||||||
|
// persona_X.session.fb — one per persona
|
||||||
|
// memory.fb — shared episodic store (existing isis.fbs)
|
||||||
|
// barriers.fb — dissociative access control
|
||||||
|
|
||||||
|
namespace isis.organism;
|
||||||
|
|
||||||
|
// ─── Shared types ────────────────────────────────────────
|
||||||
|
|
||||||
|
enum Precision : byte {
|
||||||
|
F32 = 0,
|
||||||
|
F16 = 1,
|
||||||
|
I8 = 2,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Dense matrix stored in chosen precision.
|
||||||
|
/// All weight matrices use this: synapses, NLM stages, projectors.
|
||||||
|
table Matrix {
|
||||||
|
rows: uint32;
|
||||||
|
cols: uint32;
|
||||||
|
precision: Precision;
|
||||||
|
f32_data: [float];
|
||||||
|
f16_data: [uint16];
|
||||||
|
i8_data: [int8];
|
||||||
|
i8_scale: float;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Bias vector.
|
||||||
|
table Bias {
|
||||||
|
precision: Precision;
|
||||||
|
f32_data: [float];
|
||||||
|
f16_data: [uint16];
|
||||||
|
i8_data: [int8];
|
||||||
|
i8_scale: float;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ─── Region definition ───────────────────────────────────
|
||||||
|
|
||||||
|
/// Per-neuron MLP weights within a region.
|
||||||
|
table NlmStage {
|
||||||
|
weights: Matrix; // [n_neurons × out_per × in_per]
|
||||||
|
biases: Bias; // [n_neurons × out_per]
|
||||||
|
n_neurons: uint32;
|
||||||
|
in_per: uint32;
|
||||||
|
out_per: uint32;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A brain region's learned weights.
|
||||||
|
table RegionWeights {
|
||||||
|
name: string; // "input", "attention", "output", "motor", etc.
|
||||||
|
n_neurons: uint32;
|
||||||
|
memory_length: uint32;
|
||||||
|
inhibitory_fraction: float;
|
||||||
|
inhibitory_mask: [bool];
|
||||||
|
nlm_stage1: NlmStage;
|
||||||
|
nlm_stage2: NlmStage; // null if nlm_depth < 2
|
||||||
|
start_trace: [float]; // initial trace state
|
||||||
|
start_activated: [float]; // initial activation
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Inter-region synapse weights.
|
||||||
|
table SynapseWeights {
|
||||||
|
from_region: string;
|
||||||
|
to_region: string;
|
||||||
|
weight: Matrix; // [out_dim*2 × in_dim] (×2 for GLU)
|
||||||
|
bias: Bias; // [out_dim*2]
|
||||||
|
}
|
||||||
|
|
||||||
|
// ─── Weights file (shared across personas) ───────────────
|
||||||
|
|
||||||
|
table OrganismWeights {
|
||||||
|
version: uint32;
|
||||||
|
|
||||||
|
// Architecture config
|
||||||
|
iterations: uint32;
|
||||||
|
d_model: uint32;
|
||||||
|
d_input: uint32;
|
||||||
|
n_sync_out: uint32;
|
||||||
|
n_sync_action: uint32;
|
||||||
|
motor_threshold: float;
|
||||||
|
|
||||||
|
// Brain regions (variable count — not hardcoded to 8)
|
||||||
|
regions: [RegionWeights];
|
||||||
|
|
||||||
|
// Inter-region synapses (variable count)
|
||||||
|
synapses: [SynapseWeights];
|
||||||
|
|
||||||
|
// Projectors
|
||||||
|
global_projector: Matrix;
|
||||||
|
global_projector_bias: Bias;
|
||||||
|
output_projector: Matrix;
|
||||||
|
output_projector_bias: Bias;
|
||||||
|
logit_projector: Matrix;
|
||||||
|
logit_projector_bias: Bias;
|
||||||
|
|
||||||
|
// Sync pair topology (deterministic from seed, but stored for portability)
|
||||||
|
sync_out_left: [uint32];
|
||||||
|
sync_out_right: [uint32];
|
||||||
|
sync_out_decay: [float];
|
||||||
|
sync_action_left: [uint32];
|
||||||
|
sync_action_right: [uint32];
|
||||||
|
sync_action_decay: [float];
|
||||||
|
|
||||||
|
// Position predictor weights (optional)
|
||||||
|
position_predictor_gate: Matrix;
|
||||||
|
position_predictor_content: Matrix;
|
||||||
|
position_predictor_final: Matrix;
|
||||||
|
|
||||||
|
// Organism-level learned parameters
|
||||||
|
embeddings: Matrix; // [vocab_size × embed_dim]
|
||||||
|
sensory_layers: [Matrix]; // sensory MLP weight matrices
|
||||||
|
sensory_biases: [Bias];
|
||||||
|
output_proj: Matrix; // [vocab_size × n_sync_out]
|
||||||
|
output_proj_bias: Bias;
|
||||||
|
|
||||||
|
// Metadata
|
||||||
|
vocab_size: uint32;
|
||||||
|
embed_dim: uint32;
|
||||||
|
sensory_depth: uint32;
|
||||||
|
context_len: uint32;
|
||||||
|
tokens_seen: uint64;
|
||||||
|
sleep_cycles: uint64;
|
||||||
|
created_at: float64;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ─── Session file (per-persona) ──────────────────────────
|
||||||
|
|
||||||
|
/// Neuromodulator baseline — each persona's emotional temperament.
|
||||||
|
table NeuromodBaseline {
|
||||||
|
dopamine: float;
|
||||||
|
serotonin: float;
|
||||||
|
norepinephrine: float;
|
||||||
|
acetylcholine: float;
|
||||||
|
curiosity: float;
|
||||||
|
anxiety: float;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Hippocampal content-addressable memory state.
|
||||||
|
table HippocampalState {
|
||||||
|
capacity: uint32;
|
||||||
|
key_dim: uint32;
|
||||||
|
value_dim: uint32;
|
||||||
|
keys: [float]; // [capacity × key_dim]
|
||||||
|
values: [float]; // [capacity × value_dim]
|
||||||
|
strengths: [float]; // [capacity]
|
||||||
|
write_ptr: uint32;
|
||||||
|
count: uint32;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Per-region mutable state (noise, usefulness tracking).
|
||||||
|
table RegionState {
|
||||||
|
region_name: string;
|
||||||
|
noise_scale: [float]; // [n_neurons]
|
||||||
|
usefulness_ema: [float]; // [n_neurons]
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Hebbian plasticity state per region.
|
||||||
|
table HebbianState {
|
||||||
|
region_name: string;
|
||||||
|
running_mean: [float]; // [n_neurons]
|
||||||
|
baseline_mean: [float];
|
||||||
|
baseline_var: [float];
|
||||||
|
calibrated: bool;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Replay buffer entry for prioritized consolidation.
|
||||||
|
table ReplayEntry {
|
||||||
|
observation: [float];
|
||||||
|
surprise: float;
|
||||||
|
timestamp: uint64;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A persona's complete session state.
|
||||||
|
/// This is everything that makes one persona different from another
|
||||||
|
/// running on the same shared weights.
|
||||||
|
table PersonaSession {
|
||||||
|
version: uint32;
|
||||||
|
|
||||||
|
// Identity
|
||||||
|
name: string; // persona name
|
||||||
|
created_at: float64;
|
||||||
|
|
||||||
|
// Emotional temperament
|
||||||
|
neuromod_baseline: NeuromodBaseline;
|
||||||
|
|
||||||
|
// Episodic memory (per-persona hippocampal state)
|
||||||
|
hippo: HippocampalState;
|
||||||
|
hippo_retrieval: [float];
|
||||||
|
|
||||||
|
// Basal ganglia learning state
|
||||||
|
bg_eligibility: [float];
|
||||||
|
bg_da_baseline: float;
|
||||||
|
bg_weight_delta: [float];
|
||||||
|
|
||||||
|
// Per-region mutable state
|
||||||
|
region_states: [RegionState];
|
||||||
|
|
||||||
|
// Per-region Hebbian state
|
||||||
|
hebbian_states: [HebbianState];
|
||||||
|
|
||||||
|
// Sleep consolidation traces (transient, may be empty)
|
||||||
|
sleep_trace_count: uint32;
|
||||||
|
|
||||||
|
// Replay buffer
|
||||||
|
replay_entries: [ReplayEntry];
|
||||||
|
replay_capacity: uint32;
|
||||||
|
|
||||||
|
// Training state
|
||||||
|
hebbian_enabled: bool;
|
||||||
|
tokens_processed: uint64;
|
||||||
|
gamma_cycles_total: uint64;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ─── Dissociative barriers ───────────────────────────────
|
||||||
|
|
||||||
|
/// Per-persona access control over shared episodic memory.
|
||||||
|
/// Each persona can see different episodes with different strength.
|
||||||
|
table EpisodeAccess {
|
||||||
|
episode_index: uint32; // index into MemoryBank.alters[].episodes[]
|
||||||
|
access_level: float; // 0.0 = walled off, 1.0 = full access
|
||||||
|
}
|
||||||
|
|
||||||
|
table PersonaBarrier {
|
||||||
|
persona_name: string;
|
||||||
|
episode_access: [EpisodeAccess];
|
||||||
|
}
|
||||||
|
|
||||||
|
table DissociativeBarriers {
|
||||||
|
version: uint32;
|
||||||
|
barriers: [PersonaBarrier];
|
||||||
|
// Global integration level: 0.0 = fully dissociated, 1.0 = fully integrated
|
||||||
|
integration_level: float;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ─── Root types ──────────────────────────────────────────
|
||||||
|
|
||||||
|
root_type OrganismWeights;
|
||||||
|
file_identifier "orgw";
|
||||||
|
file_extension "weights.fb";
|
||||||
|
|
||||||
|
// For session files, use PersonaSession as root:
|
||||||
|
// root_type PersonaSession;
|
||||||
|
// file_identifier "orgs";
|
||||||
|
// file_extension "session.fb";
|
||||||
|
|
||||||
|
// For barrier files, use DissociativeBarriers as root:
|
||||||
|
// root_type DissociativeBarriers;
|
||||||
|
// file_identifier "orgb";
|
||||||
|
// file_extension "barriers.fb";
|
||||||
Reference in New Issue
Block a user