Metadata-Version: 2.4
Name: sketchssm
Version: 0.1.1
Summary: SketchSSM decode kernels and offline calibration
License-Expression: Apache-2.0
Project-URL: Homepage, https://github.com/SNU-ARC/SketchSSM
Project-URL: Repository, https://github.com/SNU-ARC/SketchSSM
Requires-Python: >=3.10
Description-Content-Type: text/markdown
License-File: LICENSE
Requires-Dist: numpy>=1.26
Requires-Dist: torch>=2.6
Provides-Extra: calibration
Requires-Dist: PyYAML>=6; extra == "calibration"
Requires-Dist: transformers; extra == "calibration"
Requires-Dist: datasets; extra == "calibration"
Requires-Dist: safetensors; extra == "calibration"
Requires-Dist: huggingface-hub; extra == "calibration"
Provides-Extra: collection
Requires-Dist: sketchssm[calibration]; extra == "collection"
Provides-Extra: test
Requires-Dist: pytest; extra == "test"
Dynamic: license-file

# 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 size, which
sets the traffic reduction, is chosen in the serving configuration; the
sketching matrix is computed once per model by offline calibration.

<p align="center">
  <img src="assets/sketchssm-overview.png" alt="SketchSSM over a window" width="100%">
</p>

*SketchSSM over a window. At a flush step, SketchSSM updates the full state
S<sub>0</sub> and refreshes the sketch U. At non-flush steps, it multiplies U by
query-dependent coefficients c<sub>t</sub> to reconstruct the output; S<sub>0</sub>
is not accessed between state updates.*

[Paper](https://arxiv.org/abs/2609.33051)

This repository contains:

- [`vllm/`](https://github.com/SNU-ARC/SketchSSM/blob/main/vllm/README.md): vLLM v0.30.0 with SketchSSM.
- [`sketchssm/kernels/`](https://github.com/SNU-ARC/SketchSSM/blob/main/sketchssm/kernels/README.md): the SketchSSM CUDA decode kernels.
- [`sketchssm/calibration/`](https://github.com/SNU-ARC/SketchSSM/blob/main/sketchssm/calibration/README.md): the offline calibration.
- [`benchmarks/`](https://github.com/SNU-ARC/SketchSSM/blob/main/benchmarks/README.md): per-layer decode latency in vLLM.

## Installation

```bash
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](https://huggingface.co/ominn).
To calibrate a new model yourself, see [`sketchssm/calibration/`](https://github.com/SNU-ARC/SketchSSM/blob/main/sketchssm/calibration/README.md).

<table>
<tr><th><sub>Model</sub></th><th><sub>Weights</sub></th><th><sub>Calibration</sub></th></tr>
<tr><td><sub>Nemotron Nano 9B v2-BF16</sub></td><td><sub><a href="https://huggingface.co/nvidia/NVIDIA-Nemotron-Nano-9B-v2">nvidia/NVIDIA-Nemotron-Nano-9B-v2</a></sub></td><td><sub><a href="https://huggingface.co/ominn/SketchSSM-Nemotron-Nano-9B-v2-BF16">ominn/SketchSSM-Nemotron-Nano-9B-v2-BF16</a></sub></td></tr>
<tr><td><sub>Nemotron 3 Super-NVFP4</sub></td><td><sub><a href="https://huggingface.co/nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-NVFP4">nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-NVFP4</a></sub></td><td><sub><a href="https://huggingface.co/ominn/SketchSSM-Nemotron-3-Super-NVFP4">ominn/SketchSSM-Nemotron-3-Super-NVFP4</a></sub></td></tr>
<tr><td><sub>Qwen3.8 Flash-Next-NVFP4</sub></td><td><sub><a href="https://huggingface.co/RadixArk/Qwen3.8-Flash-Next-NVFP4">RadixArk/Qwen3.8-Flash-Next-NVFP4</a></sub></td><td><sub><a href="https://huggingface.co/ominn/SketchSSM-Qwen3.8-Flash-Next-NVFP4">ominn/SketchSSM-Qwen3.8-Flash-Next-NVFP4</a></sub></td></tr>
<tr><td><sub>GLM 5.3 Flash-NVFP4</sub></td><td><sub><a href="https://huggingface.co/RedHatAI/GLM-5.3-Flash-NVFP4">RedHatAI/GLM-5.3-Flash-NVFP4</a></sub></td><td><sub><a href="https://huggingface.co/ominn/SketchSSM-GLM-5.3-Flash-NVFP4">ominn/SketchSSM-GLM-5.3-Flash-NVFP4</a></sub></td></tr>
<tr><td><sub>Qwen3.5 9B-BF16</sub></td><td><sub><a href="https://huggingface.co/Qwen/Qwen3.5-9B">Qwen/Qwen3.5-9B</a></sub></td><td><sub><a href="https://huggingface.co/ominn/SketchSSM-Qwen3.5-9B-BF16">ominn/SketchSSM-Qwen3.5-9B-BF16</a></sub></td></tr>
</table>

<sub>See [all calibrations](https://github.com/SNU-ARC/SketchSSM/blob/main/sketchssm/calibration/README.md#portable-calibration-file).</sub>

Enable SketchSSM with:

1. `--sketchssm`: the calibration, a Hub repo id or a local file.
2. `--sketchssm-mean-rank`: the rank budget, the mean sketch rank per state head
   (default 10: about 10x less state traffic at standard decode accuracy in the paper).

```bash
# Calibration from the Hub
vllm serve nvidia/NVIDIA-Nemotron-Nano-9B-v2 --trust-remote-code \
  --sketchssm ominn/SketchSSM-Nemotron-Nano-9B-v2-BF16 --sketchssm-mean-rank 10 \
  --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 --sketchssm-mean-rank 10 \
  --mamba-ssm-cache-dtype float32 --no-enable-prefix-caching
```

See [`vllm/README.md`](https://github.com/SNU-ARC/SketchSSM/blob/main/vllm/README.md) for all options.

## Benchmarks

See [`benchmarks/`](https://github.com/SNU-ARC/SketchSSM/blob/main/benchmarks/README.md) to measure the per-layer recurrent
decode latency of a model in vLLM, and [`sketchssm/kernels/`](https://github.com/SNU-ARC/SketchSSM/blob/main/sketchssm/kernels/README.md)
for kernel tuning and microbenchmarks.

## Citation

If you use SketchSSM in your research, please cite:

```bibtex
@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](https://github.com/SNU-ARC/SketchSSM/blob/main/LICENSE). The
[`vllm/`](https://github.com/SNU-ARC/SketchSSM/blob/main/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.
