DEV Community

Cover image for CoMem Explained — From Paper to Working Code in 10 Minutes
Chaeyeon Mia Lee
Chaeyeon Mia Lee

Posted on

CoMem Explained — From Paper to Working Code in 10 Minutes

You've hit the wall: your long-context LLM pipeline eats 89 GB of VRAM for 128k tokens, your RAG system loses the thread of a long document, and every compression approach you try trades accuracy for memory. There's a new paper that reframes the whole problem — and the fix is surprisingly elegant.

TL;DR

  • What: CoMem (Comprehension Memory) — arxiv:2607.28263, published July 30 2026
  • Key insight: Lower transformer layers handle semantic understanding; upper layers handle prediction. Cache at the boundary.
  • How it works: Store residual stream states at an intermediate "split layer," retrieve top-k relevant states per query, recompute only the upper layers
  • Results: 18.26 GB vs 89.36 GB at 128k tokens; 7.83× prefill speedup; RULER score 97.05; LoCoMo +10.6% over KV-Direct baseline
  • Model: Qwen3-8B (only tested architecture so far)

The Problem

Standard KV caching stores key-value pairs for every layer across the entire context:

Memory = 2 × n_layers × seq_len × n_heads × head_dim × dtype_bytes
Enter fullscreen mode Exit fullscreen mode

For a 32-layer model, 128k tokens, FP16: that's ~85–90 GB. One GPU. Gone.

RAG alternatives chunk and re-embed context — but retrieval at the embedding level discards the structural relationships the model built during the original read. The result: disjointed responses, hallucinations about inter-paragraph relationships, the "Lost in the Middle" failure mode.

CoMem asks: if we've already encoded a document through the lower layers, why re-encode it?


How It Works

The Depth Division of Labor

The paper's central empirical claim: transformers have a functional hierarchy.

  • Lower layers (1 to ~L/2): Build semantic representations — syntax, entities, cross-token relationships
  • Upper layers (~L/2 to L): Specialize those representations for next-token prediction

This means the "understanding" phase is mostly complete by the middle of the network. The upper layers are doing task-specific translation, not comprehension.

The Three Phases

Write (offline, once per document):
Run context tokens through layers 1 to L* (split layer). Cache the residual stream state h[L*] for each token. That's it — no upper layers needed.

Retrieve (at query time):
Run the query through layers 1 to L* to get its intermediate state. Cosine-similarity search over the cache returns top-k most relevant context states.

Recompute (at query time):
Feed [retrieved_k_states + query_states] through layers L*+1 to L. The upper layers now compute over semantically relevant context — without ever re-running the lower layers.

Standard inference: O(n × L) per query (full context, all layers)
CoMem inference:    O(n × L*) once + O(k × (L - L*)) per query
                    where k << n
Enter fullscreen mode Exit fullscreen mode

Show Me The Code

import torch
import torch.nn.functional as F

class CoMemCache:
    """
    Caches intermediate layer states for CoMem inference.

    Memory footprint: O(n * d_model) — no layer dimension.
    Compare: KV cache is O(n * n_layers * d_kv).
    """

    def __init__(self, topk: int = 64):
        self.topk = topk
        self._keys = []   # normalized (d,) tensors for similarity search
        self._vals = []   # raw (d,) tensors to feed into upper layers

    def write(self, split_layer_output: torch.Tensor) -> None:
        """
        Store all token states from split layer output.
        Call once per context document.

        split_layer_output: (seq_len, d_model)
        """
        for i in range(split_layer_output.shape[0]):
            v = split_layer_output[i].detach().cpu()
            self._vals.append(v)
            self._keys.append(F.normalize(v, dim=-1))

    def retrieve(self, query_split_state: torch.Tensor) -> torch.Tensor:
        """
        Get top-k context states relevant to this query.
        Fixed cost regardless of how many tokens are cached.

        query_split_state: (d_model,)
        returns: (k, d_model)
        """
        if len(self._keys) == 0:
            return query_split_state.new_zeros(0, query_split_state.shape[0])

        keys = torch.stack(self._keys).to(query_split_state.device)
        q = F.normalize(query_split_state, dim=-1)

        scores = keys @ q                                    # (n,)
        k = min(self.topk, len(self._vals))
        _, idx = torch.topk(scores, k=k)

        vals = torch.stack(self._vals).to(query_split_state.device)
        return vals[idx]                                     # (k, d)

    def memory_mb(self, d_model: int = 4096) -> float:
        """Approximate memory usage in MB (fp16)."""
        return len(self._vals) * d_model * 2 / 1e6


def run_upper_layers(
    layers: list,            # model.layers[split+1:]
    query_hidden: torch.Tensor,      # (q_len, d)
    retrieved: torch.Tensor,         # (k, d)
) -> torch.Tensor:
    """
    Prepend retrieved context states and run upper layers on combined sequence.
    Returns output for query tokens only.
    """
    if retrieved.shape[0] > 0:
        hidden = torch.cat([retrieved, query_hidden], dim=0)  # (k+q, d)
    else:
        hidden = query_hidden

    for layer in layers:
        hidden = layer(hidden)

    # Slice off the query portion
    return hidden[-query_hidden.shape[0]:]


# --- Quick memory sanity check ---
cache = CoMemCache(topk=64)
seq_len, d_model = 128_000, 4096

# Simulate caching 128k tokens
fake_states = torch.randn(min(100, seq_len), d_model)  # sample
cache.write(fake_states)

approx_mb = seq_len * d_model * 2 / 1e6  # full 128k
print(f"CoMem cache (128k tokens): {approx_mb:.0f} MB = {approx_mb/1000:.2f} GB")
# → CoMem cache (128k tokens): 1024 MB = 1.02 GB (just states)
# Paper reports 18.26 GB total including model weights + cache + activations
# vs 89.36 GB for standard KV cache setup
Enter fullscreen mode Exit fullscreen mode

Training Adaptation (LoRA Self-Distillation)

Pre-trained models aren't optimized for receiving retrieved intermediate states. The paper adapts upper layers using LoRA:

from peft import get_peft_model, LoraConfig, TaskType
import torch.nn as nn

def setup_comem_lora(model, split_layer: int, lora_rank: int = 16):
    """
    Apply LoRA to upper layers only for CoMem adaptation.
    Lower layers stay frozen — we only need to adapt how
    upper layers process retrieved intermediate states.
    """
    # Freeze everything first
    for param in model.parameters():
        param.requires_grad = False

    # Apply LoRA config to upper layers
    lora_config = LoraConfig(
        task_type=TaskType.CAUSAL_LM,
        r=lora_rank,
        lora_alpha=32,
        target_modules=["q_proj", "v_proj"],
        # Only apply to layers above split point
        layers_to_transform=list(range(split_layer + 1, model.config.num_hidden_layers)),
    )

    model = get_peft_model(model, lora_config)

    print(f"Trainable params: {sum(p.numel() for p in model.parameters() if p.requires_grad):,}")
    return model


def self_distillation_loss(
    student_logits: torch.Tensor,  # (seq, vocab) — CoMem output
    teacher_logits: torch.Tensor,  # (seq, vocab) — full-context output
    temperature: float = 1.0,
) -> torch.Tensor:
    """
    KL divergence between student (CoMem) and teacher (full KV) distributions.
    No ground-truth labels needed — the frozen full model supervises itself.
    """
    student_log_probs = F.log_softmax(student_logits / temperature, dim=-1)
    teacher_probs = F.softmax(teacher_logits / temperature, dim=-1)
    return F.kl_div(student_log_probs, teacher_probs, reduction="batchmean")
Enter fullscreen mode Exit fullscreen mode

Benchmark Results

Metric CoMem Baseline Delta
RULER score 97.05
LoCoMo score 38.27 34.59 (KV-Direct) +3.68 (+10.6%)
Memory @ 128k 18.26 GB 89.36 GB 4.9× less
Prefill speed @ 128k 7.83× faster

Hardware: NVIDIA H20. Model: Qwen3-8B with continued training.

The RULER score is notable — 97.05 is near-ceiling performance. Standard RAG systems typically score in the 70–80% range on RULER because retrieved chunks lose inter-document context.


Gotchas & Limitations

Only tested on Qwen3-8B. The layer specialization assumption needs verification on Llama-3, Mistral, MoE architectures (Mixtral, DeepSeek). Different models may have different functional split points.

L* selection is manual. The paper uses L* ≈ L/2 but doesn't provide a principled method for finding the optimal split. You'll need to experiment.

Cosine similarity retrieval has blind spots. Multi-hop reasoning often requires context that's semantically distant from the query but logically necessary. Pure cosine similarity misses these cases. A hybrid dense+sparse retrieval would be stronger.

LoRA training required. Off-the-shelf models aren't adapted for receiving retrieved intermediate states — you need the self-distillation fine-tuning step. The paper doesn't report how much performance degrades without it.

Generation quality not evaluated. RULER and LoCoMo are retrieval-heavy. Open-ended generation quality, hallucination rate, and coherence in long documents aren't assessed.


🚀 Try It Today

The paper is at arxiv:2607.28263. While the official code isn't out yet, you can start experimenting with the concept:

  1. Profile your model's layer specialization: Measure representational similarity (CKA) across layers on your target model — find where semantics stabilize.

  2. Prototype the cache: The CoMemCache class above is a drop-in starting point. Hook it into your model's forward pass at the split layer.

  3. Test without LoRA first: Grab the upper-layer outputs without adaptation. Measure quality degradation — this tells you whether LoRA training is strictly necessary for your use case.

  4. Benchmark your target task: Before fine-tuning, verify that RULER/LoCoMo improvements generalize to your domain.

# One-liner to estimate whether CoMem is worth it for your model
def should_use_comem(current_vram_gb: float, context_tokens: int) -> str:
    # CoMem approximation: ~d_model * seq_len * 2 bytes + model weights
    # vs full KV: n_layers * d_kv * seq_len * 2 * 2 bytes + model weights
    kv_overhead_gb = 32 * 128 * context_tokens * 2 * 2 / 1e9  # 32L, 128 d_kv
    comem_overhead_gb = 4096 * context_tokens * 2 / 1e9        # just d_model

    if current_vram_gb + kv_overhead_gb > 80:  # A100/H100/H20 limit
        return f"YES — KV cache needs {kv_overhead_gb:.1f}GB, CoMem needs {comem_overhead_gb:.1f}GB"
    return f"Optional — KV cache needs {kv_overhead_gb:.1f}GB (within 80GB budget)"

print(should_use_comem(current_vram_gb=20, context_tokens=128_000))
# → YES — KV cache needs 85.9GB, CoMem needs 1.0GB
Enter fullscreen mode Exit fullscreen mode

What's your experience with long-context inference? Drop a comment — especially if you've tried alternative approaches like StreamingLLM or KV compression in production.


Sources

Top comments (0)