SketchSSM: Write to the Full State, Read from a Compact Sketch (Paper)
Overview
Hybrid-attention models replace most softmax attention layers with linear attention (Mamba-2, Gated DeltaNet, KDA), reducing KV-cache growth and enabling larger decode batches, where recurrent-state access becomes a major bottleneck. Reducing the state size by quantization or pruning cuts this traffic, but its approximation errors propagate and accumulate through subsequent decode steps. To reduce state-update traffic, ReplaySSM buffers keys and values over a window of W steps and applies their accumulated updates to the full state once per window. However, each new query still requires a full-state read, even though the state remains unchanged between state updates.
SketchSSM keeps the full-state updates and approximates the reads:
- Flush step (every W steps): read the full state S0 once, apply the buffered updates from the ring buffer, compute the compact sketch U from the updated state, and write back the state and the sketch.
- Non-flush steps: combine the sketch U with query-dependent coefficients ct to reconstruct the output, without reading S0.
The sketching matrix is computed once per model by offline calibration. The sketch size (the mean sketch rank per state head) is chosen as a serving configuration and sets the trade-off between traffic reduction and accuracy. Across Mamba-2, GDN and KDA models, SketchSSM reduces state-access traffic by about 10x while largely preserving accuracy.
This repository contains:
vllm/: vLLM v0.30.0 with SketchSSM.sketchssm/kernels/: the SketchSSM CUDA decode kernels.sketchssm/calibration/: the offline calibration.benchmarks/: per-layer decode latency in vLLM.
Installation
git clone https://github.com/SNU-ARC/SketchSSM.git
cd SketchSSM/vllm
VLLM_USE_PRECOMPILED=1 python -m pip install -e .
python -m pip install sketchssm # CUDA kernels; without it vLLM uses its Triton kernels
Quick start: serve with vLLM
SketchSSM needs a calibration file (calibration.pt) made by offline
calibration. Calibrations for the models below are on the Hugging Face Hub.
To calibrate a new model yourself, see sketchssm/calibration/.
| Model | Weights | Calibration |
|---|---|---|
| Nemotron Nano 9B v2-BF16 | nvidia/NVIDIA-Nemotron-Nano-9B-v2 | ominn/SketchSSM-Nemotron-Nano-9B-v2-BF16 |
| Nemotron 3 Super-NVFP4 | nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-NVFP4 | ominn/SketchSSM-Nemotron-3-Super-NVFP4 |
| Qwen3.8 Flash-Next-NVFP4 | RadixArk/Qwen3.8-Flash-Next-NVFP4 | ominn/SketchSSM-Qwen3.8-Flash-Next-NVFP4 |
| GLM 5.3 Flash-NVFP4 | RedHatAI/GLM-5.3-Flash-NVFP4 | ominn/SketchSSM-GLM-5.3-Flash-NVFP4 |
| Qwen3.5 9B-BF16 | Qwen/Qwen3.5-9B | ominn/SketchSSM-Qwen3.5-9B-BF16 |
See all calibrations.
Enable SketchSSM with --sketchssm and a calibration, a Hub repo id or a local
file:
# Calibration from the Hub
vllm serve nvidia/NVIDIA-Nemotron-Nano-9B-v2 --trust-remote-code \
--sketchssm ominn/SketchSSM-Nemotron-Nano-9B-v2-BF16 \
--mamba-ssm-cache-dtype float32 --no-enable-prefix-caching
# Local calibration file
vllm serve nvidia/NVIDIA-Nemotron-Nano-9B-v2 --trust-remote-code \
--sketchssm outputs/nano/calibration.pt \
--mamba-ssm-cache-dtype float32 --no-enable-prefix-caching
Optionally, --sketchssm-mean-rank sets the rank budget, the mean sketch rank
per state head (default 8: about 10x less state traffic at accuracy comparable to the
full-state baseline).
See vllm/README.md for all options.
Benchmarks
See benchmarks/ to measure the per-layer recurrent
decode latency of a model in vLLM, and sketchssm/kernels/
for kernel tuning and microbenchmarks.
Citation
If you use SketchSSM in your research, please cite:
@misc{kwon2026sketchssmwritestateread,
title={SketchSSM: Write to the Full State, Read from a Compact Sketch},
author={Omin Kwon and JoongWon Shin and Minseo Kim and Kurt Keutzer and Sehoon Kim and Jae W. Lee},
year={2026},
eprint={2609.33051},
archivePrefix={arXiv},
primaryClass={cs.LG},
url={https://arxiv.org/abs/2609.33051},
}
License
SketchSSM is released under the Apache License 2.0. The
vllm/ directory is a fork of vLLM, also under Apache-2.0. Calibration
files are derived from the base models' weights and are also subject to those
models' licenses.
Metadata
Release files for sketchssm 0.1.2
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| sketchssm-0.1.2.tar.gz | 216.9 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| sketchssm-0.1.2-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 478.0 kB
Release files / sketchssm-0.1.2.tar.gz
| Download URL | sketchssm-0.1.2.tar.gz |
|---|---|
| Size | 216.9 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
87befc3d8df5ab9f7378dda7d065f6b79c89bff194ca386af089ab9f77c751b8
|
|
BLAKE2b-256 checksum How to use checksums |
8605a16e4df1f897980bf591c925ca0e2c568daab2f8928bd205adbf0c70010b
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.12.14
|
Release files / sketchssm-0.1.2-py3-none-any.whl
| Download URL | sketchssm-0.1.2-py3-none-any.whl |
|---|---|
| Size | 261.0 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
7f46e9fc048cde4a95f5c45ed6d36abeb967abf0e76b9530607b35cb34b86948
|
|
BLAKE2b-256 checksum How to use checksums |
d683977702e3bdb74ebbabccec503798c6b3806588ed3801d9f70c008a96de54
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.12.14
|