epimem: One-shot gradient-free learning on frozen transformers
- paper.md: full paper (Tommi Niemi / Rotko Networks) - python/epimem.py: standalone Python reproduction - export_onnx.py: ONNX export from HuggingFace (generates model files) - results/memory_bank.json: example hidden-state vectors (896-dim) - schema/: FlatBuffer schemas for memory bank + organism - models/tokenizer/: Qwen 2.5 tokenizer files Run: pip install transformers torch && python python/epimem.py (Downloads Qwen 2.5 automatically from HuggingFace)
This commit is contained in:
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()
|
||||
Reference in New Issue
Block a user