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