Files
cri/python/epimem.py
Tommi Niemi 176164815b 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
2026-04-05 02:16:29 +07:00

362 lines
14 KiB
Python

#!/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()