Skip to main content ホーム クリエイター adu2021 skillxiv gradmem-context-memory-compression
gradmem-context-memory-compression Compress long context into compact memory tokens via iterative gradient descent. Learn to write information into prefix memory without storing full KV-caches, enabling efficient long-context reasoning and retrieval.
インストールへ移動 Skills Marketplace コミュニティが作成したAIスキルを発見・探索
Codex または Claude でインストール この Prompt をコピーして Codex、Claude、または他のアシスタントに貼り付けると、Skill ページを確認してインストールできます。
直接コマンドでは確認用 Prompt が省略されます。実行前にソースを確認してください。
npx skills add https://github.com/ADu2021/skillXiv --skill gradmem-context-memory-compressionコマンドは1行のまま表示されます。コピー前に横へスクロールして全体を確認してください。
ローカルで確認しますか?SkillsMP が現在取得できるファイルをダウンロードできます。
Zipをダウンロード ダウンロード中... このリポジトリの他の Skills meaningful-kebab-case-name Convert arXiv papers into ready-to-use agent skills using category-aware extraction. First classifies the paper into one or more of 11 research categories, then applies a specialized extraction pipeline for each category — because different types of papers produce different types of usable knowledge. A single paper can yield multiple skills if it spans categories. Use this skill whenever the user wants to turn a paper into a skill, extract practical techniques from research, build a skill library from papers, convert arXiv papers into reusable agent instructions, or batch-process multiple papers into skills. Also trigger when someone asks about extracting actionable knowledge from papers, making research practical for LLM agents, or systematically converting academic contributions into structured agent capabilities.
action-quantization-behavior-cloning Establish regret bounds for behavior cloning with discretized actions combining statistical error and quantization error terms. Prove smoothness requirements for safe quantizer design, show that learning-based quantizers fail these requirements, and propose model-based augmentation to reduce error dependence from H² to H.
adaptive-lora-personalized-ranks Dynamically allocate LoRA ranks per-layer during fine-tuning instead of using fixed uniform ranks. Learn optimal rank for each layer and subject via variational framework with discretized exponential distribution, reducing memory footprint while maintaining fidelity and text-alignment.
name gradmem-context-memory-compression title GradMem: Learning to Write Context into Memory with Test-Time Gradient Descent version 0.0.2 engine skillxiv-v0.0.2-claude-opus-4.6 license MIT url https://arxiv.org/abs/2603.13875 keywords ["Memory Compression","Context Storage","Gradient Descent","Long-Context","Test-Time Optimization"] description Compress long context into compact memory tokens via iterative gradient descent. Learn to write information into prefix memory without storing full KV-caches, enabling efficient long-context reasoning and retrieval.
GradMem: Test-Time Gradient Descent for Context Memory Compression
Large language models struggle with long contexts due to massive KV-cache overhead. GradMem enables models to compress context into compact memory tokens via test-time gradient descent. Rather than storing full cache for each query, the model learns to write information into a small set of learnable prefix tokens by optimizing a self-supervised reconstruction loss. This approach is particularly effective for context removal scenarios where the model must answer questions without accessing the original long context at inference time.
The key innovation is using gradient updates on memory tokens (with frozen model weights) rather than forward-only writes, enabling iterative error correction and much better context compression.
Core Concept
GradMem operates through three phases:
Context Ingestion — Process long context once to initialize memory tokens
Gradient-Based Writes — Iteratively update memory tokens using gradient descent on reconstruction loss
Query Answering — Use compressed memory (without original context) to answer questions
This contrasts with forward-only approaches which write information in a single forward pass. By optimizing iteratively, GradMem can pack more information into the same number of tokens.
Architecture Overview
Memory Token Initializer — Create learnable prefix tokens initialized from context
Reconstruction Loss — Self-supervised objective: can we reconstruct context from memory?
Gradient-Based Writer — Optimize memory tokens via SGD on reconstruction loss
Query Encoder — Generate question embeddings in the same space as memory
Memory-Only Reader — Answer questions using only memory tokens (no original context)
Scaling Mechanism — Multiple gradient steps improve capacity linearly
Implementation Steps
Start by designing the memory tokens and initializing them from context.
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.optim import SGD
class ContextMemoryCompressor :
"""Compress long context into learnable memory tokens."""
( ):
.hidden_dim = hidden_dim
.num_memory_tokens = num_memory_tokens
.context_seq_len = context_seq_len
.memory_tokens = nn.Parameter(
torch.randn(num_memory_tokens, hidden_dim) *
)
.context_encoder = nn.TransformerEncoder(
nn.TransformerEncoderLayer(d_model=hidden_dim, nhead= ,
dim_feedforward= ,
batch_first= ),
num_layers=
)
.context_decoder = nn.TransformerDecoder(
nn.TransformerDecoderLayer(d_model=hidden_dim, nhead= ,
dim_feedforward= ,
batch_first= ),
num_layers=
)
( ):
batch_size = context_embeddings.size( )
memory = .memory_tokens.unsqueeze( ).expand(batch_size, - , - )
cross_attn = nn.MultiheadAttention( .hidden_dim, num_heads= ,
batch_first= )
memory_initialized, _ = cross_attn(memory, context_embeddings,
context_embeddings)
memory_initialized
( ) -> torch.Tensor:
reconstructed = .context_decoder(memory, context)
loss = F.mse_loss(reconstructed, context)
loss
( ) -> torch.Tensor:
batch_size = context_embeddings.size( )
memory = .initialize_memory(context_embeddings)
memory_param = torch.nn.Parameter(memory.clone())
optimizer = SGD([memory_param], lr=learning_rate)
step (num_gradient_steps):
optimizer.zero_grad()
loss = .reconstruction_loss(memory_param, context_embeddings)
loss.backward()
optimizer.step()
( )
memory_param.detach()
( ) -> torch.Tensor:
memory = context_embeddings.mean(dim= , keepdim= ).expand(
- , .num_memory_tokens, -
)
memory
def
__init__
self, hidden_dim=768 , num_memory_tokens=32 , context_seq_len=8192
self
self
self
self
0.01
self
8
2048
True
2
self
8
2048
True
2
def
initialize_memory
self, context_embeddings: torch.Tensor
"""Initialize memory tokens from context."""
0
self
0
1
1
self
8
True
return
def
reconstruction_loss
self, memory: torch.Tensor,
context: torch.Tensor
"""Loss for reconstructing context from memory."""
self
return
def
write_to_memory_gradient_based
self, context_embeddings: torch.Tensor,
num_gradient_steps=5 ,
learning_rate=0.1
"""Iteratively optimize memory tokens via gradient descent."""
0
self
for
in
range
self
print
f" Gradient step {step+1 } : Loss = {loss.item():.4 f} "
return
def
forward_write
self, context_embeddings: torch.Tensor
"""Baseline: single forward pass to write context (for comparison)."""
1
True
1
self
1
return
Now implement the query answering mechanism that uses only memory without original context.
class MemoryOnlyReader (nn.Module):
"""Answer questions using only memory tokens."""
def __init__ (self, hidden_dim=768 , vocab_size=50257 ):
super ().__init__()
self .hidden_dim = hidden_dim
self .decoder = nn.TransformerDecoder(
nn.TransformerDecoderLayer(d_model=hidden_dim, nhead=8 ,
dim_feedforward=2048 ,
batch_first=True ),
num_layers=3
)
self .output_proj = nn.Linear(hidden_dim, vocab_size)
def forward (self, query_embeddings: torch.Tensor,
memory: torch.Tensor ) -> torch.Tensor:
"""
Args:
query_embeddings: [batch_size, query_len, hidden_dim]
memory: [batch_size, num_memory_tokens, hidden_dim]
Returns:
logits: [batch_size, query_len, vocab_size]
"""
output = self .decoder(query_embeddings, memory)
logits = self .output_proj(output)
return logits
Integrate gradient-based writing into full pipeline with benchmarking.
class GradMemModel (nn.Module):
"""Full GradMem system for context compression and QA."""
def __init__ (self, base_model, hidden_dim=768 , num_memory_tokens=32 ):
super ().__init__()
self .base_model = base_model
self .compressor = ContextMemoryCompressor(hidden_dim, num_memory_tokens)
self .memory_reader = MemoryOnlyReader(hidden_dim)
def process_context (self, context_ids: torch.Tensor,
num_gradient_steps=5 ) -> torch.Tensor:
"""Compress context into memory tokens."""
context_embeddings = self .base_model.encoder(context_ids)
memory = self .compressor.write_to_memory_gradient_based(
context_embeddings,
num_gradient_steps=num_gradient_steps
)
return memory
def answer_question (self, query_ids: torch.Tensor,
memory: torch.Tensor ) -> torch.Tensor:
"""Answer question using memory (without original context)."""
query_embeddings = self .base_model.encoder(query_ids)
logits = self .memory_reader(query_embeddings, memory)
return logits
def forward (self, context_ids: torch.Tensor, query_ids: torch.Tensor,
num_gradient_steps=5 ) -> torch.Tensor:
"""Full pipeline: compress context, then answer query."""
memory = self .process_context(context_ids, num_gradient_steps)
logits = self .answer_question(query_ids, memory)
return logits
def benchmark_memory_compression (model, test_cases, baselines=['forward_only' ] ):
"""Measure compression efficiency vs accuracy."""
results = []
for num_steps in [1 , 3 , 5 , 10 ]:
accuracies = []
reconstruction_errors = []
for case in test_cases:
context = case ['context' ]
query = case ['query' ]
reference_answer = case ['answer' ]
context_ids = tokenizer.encode(context)
context_embeddings = model.base_model.encoder(
torch.tensor([context_ids]))
memory = model.compressor.write_to_memory_gradient_based(
context_embeddings,
num_gradient_steps=num_steps
)
query_ids = tokenizer.encode(query)
logits = model.answer_question(torch.tensor([query_ids]), memory)
predicted_answer = tokenizer.decode(torch.argmax(logits, dim=-1 ))
accuracy = compute_em(predicted_answer, reference_answer)
accuracies.append(accuracy)
reconstructed = model.compressor.context_decoder(memory,
context_embeddings)
recon_error = F.mse_loss(reconstructed,
context_embeddings).item()
reconstruction_errors.append(recon_error)
avg_accuracy = sum (accuracies) / len (accuracies)
avg_recon_error = sum (reconstruction_errors) / len (reconstruction_errors)
results.append({
'num_gradient_steps' : num_steps,
'accuracy' : avg_accuracy,
'reconstruction_error' : avg_recon_error
})
print (f"Gradient steps {num_steps} : Accuracy={avg_accuracy:.1 %} , "
f"Recon Error={avg_recon_error:.4 f} " )
return results
Practical Guidance Hyperparameters and When to Use:
Number of memory tokens typically 16-64; more tokens improve capacity but increase overhead
Gradient steps 3-10 work well; each step scales capacity roughly linearly
Learning rate 0.01-0.5; lower values converge more slowly, higher values may diverge
Use when context is fixed across multiple questions (e.g., document QA, retrieval)
Particularly effective for associative retrieval and structured reasoning tasks
For streaming scenarios where context changes frequently (memory becomes stale)
When exact context retrieval is needed (compression loses information)
For very short contexts where compression overhead exceeds benefits
When latency is critical; gradient descent adds test-time cost
Memory tokens becoming misaligned with query encoder; use shared embeddings
Reconstruction loss not capturing important information; weight loss by query relevance
Gradient descent diverging; use gradient clipping and careful learning rate selection
Not accounting for information loss in compression; use reconstruction error as proxy for quality
Reference