SketchSSM
SketchSSM: Write to the Full State, Read from a Compact Sketch
SketchSSM speeds up decoding in linear-attention and state-space layers (Mamba-2, Gated DeltaNet, KDA). It keeps updating the full recurrent state, but between periodic exact flushes it reads a compact low-rank sketch of the state instead of the full state, which cuts state-read traffic. The sketch basis and per-head ranks come from a one-time offline calibration of each model.
This repository contains:
sketchssm/kernels/: the CUDA decode kernels (sketchssm.kernels), which vLLM uses when this package is installed.sketchssm/calibration/: calibrates a model and packages the result into one portablecalibration.pt.vllm/: vLLM v0.30.0 with SketchSSM decode kernels.benchmarks/: per-layer decode latency in vLLM.
Installation
Python 3.10 or newer with PyTorch. Git LFS is not needed: calibrations for serving are downloaded from the Hugging Face Hub when used.
git clone https://github.com/SNU-ARC/SketchSSM.git
cd SketchSSM
python -m pip install ".[calibration]" # kernels + calibration
(cd vllm && VLLM_USE_PRECOMPILED=1 python -m pip install -e .) # serving
Quick start: serve with vLLM
Pass a calibration (a Hugging Face repo id or a local file) and a rank budget, the mean sketch rank per state head:
vllm serve nvidia/NVIDIA-Nemotron-Nano-9B-v2 --trust-remote-code \
--sketchssm ominn/SketchSSM-Nemotron-Nano-9B-v2-BF16 --sketchssm-mean-rank 6 \
--mamba-ssm-cache-dtype float32 --no-enable-prefix-caching
The SketchSSM CUDA kernels are compiled just in time the first time a layer
shape is used, which needs the CUDA toolkit's nvcc and ninja
(pip install ninja) on PATH; without them vLLM uses the portable Triton
kernels. The per-head ranks and frames for the requested budget are derived when the
model loads. Provided calibrations:
| Model | Calibration | Collected with weights | Calibrated budgets |
|---|---|---|---|
| Nemotron Nano 9B v2 | ominn/SketchSSM-Nemotron-Nano-9B-v2-BF16 |
BF16 nvidia/NVIDIA-Nemotron-Nano-9B-v2 |
4, 6, 10, 21 |
| Nemotron 3 Super | ominn/SketchSSM-Nemotron-3-Super-NVFP4 |
NVFP4 nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-NVFP4 |
2, 3, 5, 9, 20 |
| Qwen3.8 Flash-Next | ominn/SketchSSM-Qwen3.8-Flash-Next-NVFP4 |
NVFP4 RadixArk/Qwen3.8-Flash-Next-NVFP4 |
3, 4, 7, 11, 26 |
| GLM 5.3 Flash | ominn/SketchSSM-GLM-5.3-Flash-NVFP4 |
NVFP4 RedHatAI/GLM-5.3-Flash-NVFP4 |
3, 4, 7, 12, 28 |
| Qwen3.5 9B | ominn/SketchSSM-Qwen3.5-9B-BF16 |
BF16 Qwen/Qwen3.5-9B |
3, 4, 7, 11, 26 |
A calibration matches the weights it was collected with. Serve it with those
weights, or make a new calibration for other weights. The calibrated budgets
were checked against the stored allocation tables; other budgets use the same
allocation rule. See vllm/README.md for all options and
constraints.
Make a calibration for your model
Write a configuration for your model (start from a nearby
sketchssm/calibration/example/<model>/collect.yaml, following the
new-model guide), then run the
calibration stages:
python -m sketchssm.calibration calibrate --config my_model.yaml --out outputs/my_model
generate: generate continuations from WikiText-2 prompts.covariance: replay them and collect state and query covariances.basis: fit one shared sketch basis per head group.paired: score each rank on validation text with output-error and gradient pairs.allocate: assign each head a rank or a dense fallback for every requested budget.
Add --stage <name> to run one stage or --resume to continue. Package the
result into a single portable file, then serve it locally or share it on the Hub:
python -m sketchssm.calibration package --bundle outputs/my_model --out outputs/my_model/calibration.pt
vllm serve <your-model> --sketchssm outputs/my_model/calibration.pt --sketchssm-mean-rank 6 \
--mamba-ssm-cache-dtype float32 --no-enable-prefix-caching
python -m sketchssm.calibration manifest --bundle outputs/my_model \
--calibration outputs/my_model/calibration.pt --precision bf16 --out outputs/my_model/manifest.json
hf upload <user>/<repo> outputs/my_model/calibration.pt calibration.pt # then --sketchssm <user>/<repo>
hf upload <user>/<repo> outputs/my_model/manifest.json manifest.json
manifest records the file hash, the base checkpoint and, for every calibrated
budget, that the tables and frames derived from calibration.pt equal the
bundle's.
The offline calibration guide covers the procedure, configuration, file formats and a small CPU-only example.
Reproduce the provided calibrations
The provided calibrations were collected with the checkpoints and revisions
pinned in sketchssm/calibration/example/<model>/config.yaml. Running
calibrate with that model's collect.yaml regenerates them:
python -m sketchssm.calibration calibrate \
--config sketchssm/calibration/example/nemotron_super/collect.yaml --out outputs/nemotron_super
python -m sketchssm.calibration package --bundle outputs/nemotron_super \
--out outputs/nemotron_super/calibration.pt
Recollection is not needed to try other rank budgets. Each Hub
calibration.pt already contains the sketch basis and the allocation scores,
so the per-head table and frames for any budget are derived from it directly:
hf download ominn/SketchSSM-Nemotron-3-Super-NVFP4 calibration.pt --local-dir outputs/super
python -m sketchssm.calibration export --calibration outputs/super/calibration.pt \
--mean-rank 5 --out outputs/super_g5_frames.pt
See the calibration data guide for the collection recipe and the expected results of each model.
Benchmarks
See benchmarks/ to measure the per-layer recurrent
decode latency of a model in vLLM, and vllm/README.md for
kernel 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.0
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.0.tar.gz | 201.9 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| sketchssm-0.1.0-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 443.6 kB
Release files / sketchssm-0.1.0.tar.gz
| Download URL | sketchssm-0.1.0.tar.gz |
|---|---|
| Size | 201.9 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
293a47a3689d20de27a94b190f31c4999352238f09341c822daa150b314c3b7d
|
|
BLAKE2b-256 checksum How to use checksums |
123f1a79686c930c1b2546a0de80208f3f47138c41b49029d6a54726623ba1c5
|
| 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.0-py3-none-any.whl
| Download URL | sketchssm-0.1.0-py3-none-any.whl |
|---|---|
| Size | 241.8 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
a3db100a5a9ecb187f9c5b37484edb63d7be7a11ac8f9bca9e5d7be8b8487652
|
|
BLAKE2b-256 checksum How to use checksums |
bf0d033abc8f8bf2fb5b5317b754e209bd888cb2df185fa8fea4c2161763eef2
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.12.14
|