Metadata-Version: 2.4
Name: sketchssm
Version: 0.1.2
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: Write to the Full State, Read from a Compact Sketch ([Paper](https://arxiv.org/abs/2609.33051))

## 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 S<sub>0</sub> 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
  c<sub>t</sub> to reconstruct the output, without reading S<sub>0</sub>.

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.

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

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 `--sketchssm` and a calibration, a Hub repo id or a local
file:

```bash
# 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`](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.
