SketchSSM: Write to the Full State, Read from a Compact Sketch

📅 2026-09-27
📈 Citations: 0
✨ Influential: 0
📄 PDF
🤖 AI Summary
This study addresses the severe memory bottleneck caused by full-state reads during the decoding phase of hybrid attention models. We propose a "write-full, read-sketch" mechanism that integrates linear attention with low-rank approximation. During updates, compact state sketches are precomputed and stored; at decoding, outputs are reconstructed via coefficients to circumvent full-state access. Furthermore, basis vectors are fixed offline to effectively approximate unknown future queries. This approach reduces state access traffic by approximately 10×, achieves kernel speedups of up to 7.78×, and improves throughput by 2.64× while maintaining strong accuracy.
📝 Abstract
Hybrid-attention models replace most softmax attention layers with linear attention, reducing KV-cache growth and enabling larger decode batches where recurrent state access becomes a major bottleneck. ReplaySSM amortizes state updates by buffering keys and values, but each new query still requires a full-state read even though the state remains unchanged between state updates. We observe that low-rank state-weighted query approximation accurately preserves state-read outputs. Although future queries are unknown, the basis vectors used to approximate them can be fixed offline. Based on this observation, we introduce SketchSSM, which preserves full-state updates while approximating reads. At each state update, SketchSSM reads the full state once to precompute outputs for these basis vectors, storing them in a compact sketch. Each subsequent decode step combines the sketch vectors with query-dependent coefficients to reconstruct the output without a full-state read. Across four Mamba-2-, GDN-, and KDA-based models, SketchSSM reduces state-access traffic by approximately 10x while largely preserving average accuracy across four decode benchmarks and recall on four RULER retrieval tasks. On one NVIDIA B300, linear-attention kernel speedups over the standard vLLM baseline reach 7.78x, 5.22x, and 5.20x for Mamba-2, GDN, and KDA, respectively, with up to 2.64x higher decode throughput on Nemotron 3 Super.
Problem

Research questions and friction points this paper is trying to address.

hybrid-attention models
state-access bottleneck
linear attention
KV-cache
decode throughput
Innovation

Methods, ideas, or system contributions that make the work stand out.

SketchSSM
Linear Attention
Low-rank Approximation
KV-cache Optimization
State-space Models