Titans-Style Neural Long-Term Memory for Transformers
MemoryMAG implements a MAG (Memory as a Gate) architecture based on Google's Titans paper. It augments standard transformers with neural long-term memory, enabling O(1) retrieval complexity for million-token contexts.
Instead of relying solely on attention (which scales O(n²) with context), we add a parallel memory branch at every layer:
- Write: Information stored into fixed-size MLP weights based on "surprise" (how unexpected the input is)
- Read: Retrieved via learned query vectors that resonate with stored patterns
- Gate: Learned mixing between attention and memory outputs
The model learns how to save, how to query, and when to use memory vs. attention - all through backpropagation.
┌─────────────────────────────────────────┐
│ Decoder Layer N │
├─────────────────────────────────────────┤
│ │
│ h_in ──┬── [Attention] ──→ attn_out │
│ │ │ │
│ └── [Query Proj] ──→ [NMM] ──→ ltm_out
│ │ │
│ ┌─────────────────┘ │
│ ▼ │
│ [Gate: g = σ(W·h)] │
│ │ │
│ ▼ │
│ h_out = residual + (1-g)·attn + g·ltm │
└─────────────────────────────────────────┘
- Titans Paper: "Learning to Memorize at Test Time" (Google Research, 2024)
- MIRAS Framework: Theoretical unification of sequence modeling as associative memory
- Results: Outperforms GPT-4 on long-context benchmarks, scales to 2M+ tokens
# Create virtual environment
python -m venv venv
source venv/bin/activate
# Install dependencies
pip install torch transformers einops tqdm accelerate
# For Qwen3 models
pip install transformers>=4.40.0from src import patch_qwen3_with_mag, Qwen3MAGConfig
# Configure MAG components
config = Qwen3MAGConfig(
memory_layers=2, # Depth of memory MLP
n_persistent_tokens=16, # Learned prefix tokens
chunk_size=64, # Memory update chunk size
attention_window=None, # Limit attention span to force memory use
)
# Load and patch model
model = patch_qwen3_with_mag(
model_name_or_path="Qwen/Qwen3-1.7B",
config=config,
device="auto",
dtype=torch.bfloat16,
)
# Check trainable parameters
trainable = model.count_trainable_parameters()
print(f"Trainable: {trainable:,} parameters")# Phase 1: Hash-Hop (exact match retrieval)
python data/generate_hash_hop.py \
--context_len 64000 \
--num_samples 10000 \
--output data/hash_hop_64k.jsonl
# Phase 2: Dependency Tracing (multi-hop reasoning)
python data/generate_dependency.py \
--hops 3 \
--num_samples 10000 \
--output data/dependency_3hop.jsonl
# Phase 3: Code Retrieval (synthetic)
python data/generate_code_retrieval.py \
--synthetic \
--num_samples 10000 \
--output data/code_retrieval.jsonl# Phase 1: Hash-Hop (exact retrieval)
python training/train_mag.py \
--model_name Qwen/Qwen3-1.7B \
--data_path data/hash_hop_16k.jsonl \
--output_dir checkpoints/phase1 \
--learning_rate 1e-4 \
--num_epochs 3 \
--patch_layers every_4 \
--memory_layers 1 \
--chunk_size 16 \
--n_persistent_tokens 0 \
--gate_init_bias -4.0 \
--memory_lr 1e-4 \
--memory_momentum 0.9 \
--memory_weight_decay 0.01 \
--memory_max_update_norm 0.1 \
--memory_surprise_threshold 0.0 \
--attention_window 4096 \
--gradient_checkpointing \
--optim_8bit
# Phase 2: Dependency (builds on Phase 1 checkpoint)
python training/train_mag.py \
--model_name Qwen/Qwen3-1.7B \
--data_path data/dependency_2hop.jsonl \
--output_dir checkpoints/phase2 \
--resume_from checkpoints/phase1/best \
--learning_rate 5e-5 \
--num_epochs 3
# Phase 3: Code (builds on Phase 2 checkpoint)
python training/train_mag.py \
--model_name Qwen/Qwen3-1.7B \
--data_path data/code_retrieval.jsonl \
--output_dir checkpoints/phase3 \
--resume_from checkpoints/phase2/best \
--learning_rate 3e-5 \
--num_epochs 3All phases support the same core MAG flags:
--attention_window(force memory use by limiting attention)--patch_layers(e.g.,every_4,last_1, or indices)--memory_layers,--chunk_size,--n_persistent_tokens--gate_init_bias,--gate_reg_weight,--gate_saturation_threshold,--gate_min_std--memory_lr,--memory_momentum,--memory_weight_decay--memory_max_update_norm,--memory_surprise_threshold--gradient_checkpointing,--optim_8bit,--load_in_8bit,--load_in_4bit
Checkpoints chain automatically - each phase builds on the previous phase's learned weights.
# Needle-in-haystack benchmark
python evaluation/needle_test.py \
--checkpoint checkpoints/best \
--context_lengths 2000,4000,8000,16000 \
--patch_layers every_4 \
--memory_layers 1 \
--chunk_size 16 \
--n_persistent_tokens 0 \
--attention_window 4096 \
--attn_implementation sdpa
# Code completion evaluation
python evaluation/code_completion.py \
--checkpoint checkpoints/best# Run diagnostics on checkpoint
python training/diagnostics.py \
--checkpoint checkpoints/latestMemoryMAG/
├── SPEC.md # Full technical specification
├── README.md # This file
│
├── src/
│ ├── __init__.py
│ ├── neural_memory.py # Deep MLP memory module
│ ├── query_projector.py # Query generation with layer refinement
│ ├── mag_layer.py # Augmented decoder layer
│ ├── patch_model.py # Qwen3 patching utilities
│ └── utils.py # Helper functions
│
├── data/
│ ├── generate_hash_hop.py # Phase 1 data
│ ├── generate_dependency.py # Phase 2 data
│ └── generate_code_retrieval.py # Phase 3 data
│
├── training/
│ ├── train_mag.py # Main training loop
│ ├── curriculum.py # Training phase management
│ └── diagnostics.py # Gate/memory monitoring
│
└── evaluation/
├── needle_test.py # Needle-in-haystack
└── code_completion.py # Code completion tasks
Training proceeds in phases with increasing memory pressure:
| Phase | Data | Context | Goal |
|---|---|---|---|
| 1a | Hash-Hop (short) | 8k-16k | Gates learn to open |
| 1b | Hash-Hop (long) | 32k-64k | Memory learns to persist |
| 2a | Dependency (2-hop) | 32k | Query refinement basics |
| 2b | Dependency (3-4 hop) | 64k | Deep reasoning chains |
| 3a | Code (clean) | 64k-128k | Semantic retrieval |
| 3b | Code (complex) | 128k+ | Noise resistance |
The memory is a 2+ layer MLP whose weights constitute the memory storage:
- Surprise metric: Reconstruction error drives updates
- Momentum: Captures context around surprising events
- Weight decay: Adaptive forgetting prevents saturation
Each layer's memory output feeds the next layer's query projector:
- Early layers: Broad, syntax-focused queries
- Middle layers: Relationship-focused queries
- Late layers: Intent-focused, precise queries
Learns when to use attention vs. memory:
- Starts biased toward attention (gate ≈ 0)
- Opens on tokens requiring long-range retrieval
- Per-layer, per-token decisions
You can cap attention to a fixed window while still feeding long contexts:
attention_window = 4096means attention only sees the last 4096 tokens- Memory still sees the full sequence, forcing retrieval for out-of-window facts
- Development: 128GB+ RAM, GPU with 24GB+ VRAM
- Target: Qwen3-1.7B fits on consumer GPUs
- Production: A100/H100 for full training
- SPEC.md - Complete technical specification
- Detailed architecture, training strategies, and success metrics
- Titans: Learning to Memorize at Test Time (Google Research, 2024)
- MIRAS: A Unified Framework for Sequence Modeling (Google Research, 2024)
- Google Research Blog
- Config:
max_seq_length=8192,attention_window=4096,patch_layers=every_4,memory_layers=1,chunk_size=16,n_persistent_tokens=0 - Dataset:
data/hash_hop_16k.jsonl(4,000 samples) - Batch:
batch_size=2,gradient_accumulation_steps=4,num_epochs=1 - Optim:
learning_rate=1e-4,optim_8bit,gradient_checkpointing - Result: loss trended down to ~3.5 by step 500
- Gate analysis: layers 20 and 24 show active gates (std > 0.01); earlier layers remain near-zero (expected with
gate_init_bias=-4.0) - Memory analysis: weight norms stable (~49.7) across patched layers
- Eval mode: teacher-forcing (
--eval_mode teacher_forcing) - Context length: 8192 with
attention_window=4096 - Result: 0% for checkpoint and 0% baseline (/dev/null)
- MAG-patched model with checkpoint generates normal text (no "pipe" artifacts)
- Increase training compute (more steps/epochs, larger dataset)
- Consider
memory_layers=2for capacity once stable - Revisit gate bias/regularization after longer training
- Too few optimization steps for random-hash retrieval to shape memory usage
- Gates stayed near zero in early layers, starving the memory path of gradient
- Memory config was conservative (
memory_layers=1, larger chunking), reducing update strength - Evaluation metric is strict exact-match; any near-miss counts as 0
MIT