Skip to content

unixsysdev/MemoryMAG

Folders and files

NameName
Last commit message
Last commit date

Latest commit

 

History

54 Commits
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

MemoryMAG

Titans-Style Neural Long-Term Memory for Transformers

Overview

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.

Key Innovation

Instead of relying solely on attention (which scales O(n²) with context), we add a parallel memory branch at every layer:

  1. Write: Information stored into fixed-size MLP weights based on "surprise" (how unexpected the input is)
  2. Read: Retrieved via learned query vectors that resonate with stored patterns
  3. 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.

Architecture

┌─────────────────────────────────────────┐
│           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  │
└─────────────────────────────────────────┘

Research Foundation

  • 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

Installation

# 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.0

Quick Start

1. Patch a Model with MAG

from 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")

2. Generate Training Data

# 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

3. Train MAG Components

# 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 3

All 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.

4. Evaluate

# 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

5. Monitor Training

# Run diagnostics on checkpoint
python training/diagnostics.py \
    --checkpoint checkpoints/latest

Project Structure

MemoryMAG/
├── 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 Curriculum

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

Key Components

Neural Memory Module

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

Query Refinement

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

Gate

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

Attention Window (Memory Pressure)

You can cap attention to a fixed window while still feeding long contexts:

  • attention_window = 4096 means attention only sees the last 4096 tokens
  • Memory still sees the full sequence, forcing retrieval for out-of-window facts

Hardware Requirements

  • Development: 128GB+ RAM, GPU with 24GB+ VRAM
  • Target: Qwen3-1.7B fits on consumer GPUs
  • Production: A100/H100 for full training

Documentation

  • SPEC.md - Complete technical specification
  • Detailed architecture, training strategies, and success metrics

References

  1. Titans: Learning to Memorize at Test Time (Google Research, 2024)
  2. MIRAS: A Unified Framework for Sequence Modeling (Google Research, 2024)
  3. Google Research Blog

Results (Current Branch)

Training Run (Phase 1)

  • 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

Diagnostics Snapshot

  • 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

Needle Test (Attention Window Enforced)

  • Eval mode: teacher-forcing (--eval_mode teacher_forcing)
  • Context length: 8192 with attention_window=4096
  • Result: 0% for checkpoint and 0% baseline (/dev/null)

Sanity Check (Generation)

  • MAG-patched model with checkpoint generates normal text (no "pipe" artifacts)

Next Steps

  • Increase training compute (more steps/epochs, larger dataset)
  • Consider memory_layers=2 for capacity once stable
  • Revisit gate bias/regularization after longer training

Why This Failed (Critical Notes)

  • 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

License

MIT

About

Titans-style neural long-term memory for Qwen transformers

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

No releases published

Packages

 
 
 

Contributors

Languages