Skip to content

Repository files navigation

Grafted Titans

Хирургическая имплантация explicit memory в замороженную LLM. Обучаются только memory-адаптеры (~21M параметров), базовая модель (~4B) заморожена полностью.


Результаты

Версия KV Retrieval Passkey Overall Шагов
v5 (prefill memory) 66.2% 0% 47.0% 2800
v6 (decode memory) 32.4% 100% 52.0% 3500
v7 (cached decode) В процессе В процессе
  • KV Retrieval: "Alice lives in Paris" → "Where does Alice live?" → "Paris"
  • Passkey: "The secret code is 342" + noise → "What was the secret code?" → "342"
  • Base model: Qwen3-4B-Instruct, 4 injection layers [0, 9, 18, 27]

Архитектура

Explicit Stored Hidden States + KV-Concat

Turn 1 (memorize):
  context → frozen Qwen → capture hidden states at inject layers
                          → stored_hidden = {layer_idx: tensor}

Turn 2 (retrieve + generate):
  question → frozen Qwen with InjectedAttention:
    At each inject layer:
      stored_hidden[layer] → KVConcatAdapter → mem_K, mem_V
      attention([mem_K|seq_K], [mem_V|seq_V]) → content-based retrieval

KVConcatAdapter (trainable, ~5.2M per layer):

  • input_norm (RMSNorm, init from Qwen's input_layernorm)
  • mem_to_k, mem_to_v (Linear, init from Qwen's k_proj/v_proj)
  • k_norm (RMSNorm, init from Qwen's k_norm)
  • k_pos_embed (learnable positional embeddings for token order)
  • V-gate: sigmoid(learnable_scalar), init=0.5

Decode Memory

Memory нужна не только на prefill (вопрос), но и на каждом decode шаге (генерация ответа). Критично для multi-token ответов (passkey digits).

Adapter projection кешируется на prefill и переиспользуется на decode — без пересчёта на каждом шаге.

Что не работает

  • NeuralMemory из titans-pytorch: gradient-updated MLP хранит сжатое "резюме", а не отдельные факты. cos(retrieve(A), retrieve(B)) = 0.99 — не различает запросы. Заменена на explicit stored hidden states.
  • Memory в KV cache: попытка хранить memory внутри Qwen KV cache ломает размеры cache между inject и non-inject layers.

Ключевые находки

  1. Explicit memory >> NeuralMemory: 66% vs 2.5%. Attention-based retrieval по stored hidden states работает, gradient-updated MLP — нет.

  2. Decode memory необходима для multi-token ответов: без неё passkey=0% (модель видит memory только на первом токене, остальные digit-ы угадывает).

  3. Norm matching критичен: adapter KV нормы должны совпадать с Qwen KV нормами. Без input_norm + k_norm — 15-22x mismatch → gibberish.

  4. SDPA не backprop-ит через concat'd KV с GQA: нужен eager_attention_forward на inject layers.

  5. Gate init: bias=-3.0 → vanishing gradients (sigmoid'=0.045). bias=0.0 (sigmoid=0.5) работает.

  6. Overfitting: loss 0.04 при eval 32% → модель переобучается. Early stopping с patience=5 ловит пик.


Структура

src/
  model/
    grafted_titan.py      # GraftedTitanQwen: memorize(), retrieve_and_generate(), generate()
    injected_layer.py      # InjectedAttention: cached decode memory
    memory_adapter.py      # KVConcatAdapter: projections + norms + gate + pos_embed
    neural_memory.py       # NeuralMemoryWrapper (deprecated, kept for reference)
  training/
    trainer.py             # Training loop + eval + early stopping
  data/
    synthetic_dataset.py   # KV + passkey tasks
    babilong_dataset.py    # BABILong dataset loader
  eval/
    evaluate.py            # Evaluation pipeline
configs/
  qwen3_4b_instruct.yaml  # Current config
run_train_v3.py            # Training entry point

Что обучается vs заморожено

Компонент Статус Параметры
Qwen3-4B FROZEN ~4B
KVConcatAdapter x4 TRAINABLE ~21M
Всё остальное FROZEN

Запуск

# vast.ai: A100 SXM4 40GB
pip install torch transformers datasets pyyaml titans-pytorch==0.5.3

# Тренировка с early stopping
python run_train_v3.py

# Мониторинг
grep -E "(Eval step|New best|Early stop)" train_*.txt

Открытые вопросы

  • KV vs PK trade-off: memory на decode помогает passkey (100%) но мешает KV retrieval (66% → 32%). Cached decode (v7) — текущий компромисс.
  • Overfitting: модель быстро переобучается (loss→0.04, eval stagnates). Нужны: dropout, data augmentation, или регуляризация.
  • Масштабирование контекста: текущие тесты на коротких контекстах (40-100 tokens). BABILong 4k/16k — следующий шаг.

Предыстория

Google предложил Titans (NeurIPS 2025): Neural Memory модуль с test-time learning через mini-gradient descent. Код не выложен.

Maria Sukhareva (январь 2026) сделала proof-of-concept на Qwen-2.5-0.5B: Layer 0 injection, cross-attention, 44.7% vs 34% baseline.

Этот проект: multi-layer injection на Qwen3-4B-Instruct, explicit memory вместо NeuralMemory, 66% KV retrieval accuracy.


Ссылки

About

No description, website, or topics provided.

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages