Metadata-Version: 2.4
Name: sketchssm
Version: 0.1.0
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
Requires-Dist: ninja
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 basis and
per-head ranks come from a one-time offline calibration of each model.

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

This repository contains:

- [`sketchssm/kernels/`](https://github.com/SNU-ARC/SketchSSM/blob/main/sketchssm/kernels/README.md): the CUDA decode kernels
  (`sketchssm.kernels`), which vLLM uses when this package is installed.
- [`sketchssm/calibration/`](https://github.com/SNU-ARC/SketchSSM/blob/main/sketchssm/calibration/README.md): calibrates a model
  and packages the result into one portable `calibration.pt`.
- [`vllm/`](https://github.com/SNU-ARC/SketchSSM/blob/main/vllm/README.md): vLLM v0.30.0 with SketchSSM decode kernels.
- [`benchmarks/`](https://github.com/SNU-ARC/SketchSSM/blob/main/benchmarks/README.md): 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.

```bash
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:

```bash
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`](https://huggingface.co/ominn/SketchSSM-Nemotron-Nano-9B-v2-BF16) | BF16 [`nvidia/NVIDIA-Nemotron-Nano-9B-v2`](https://huggingface.co/nvidia/NVIDIA-Nemotron-Nano-9B-v2) | 4, 6, 10, 21 |
| Nemotron 3 Super | [`ominn/SketchSSM-Nemotron-3-Super-NVFP4`](https://huggingface.co/ominn/SketchSSM-Nemotron-3-Super-NVFP4) | NVFP4 [`nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-NVFP4`](https://huggingface.co/nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-NVFP4) | 2, 3, 5, 9, 20 |
| Qwen3.8 Flash-Next | [`ominn/SketchSSM-Qwen3.8-Flash-Next-NVFP4`](https://huggingface.co/ominn/SketchSSM-Qwen3.8-Flash-Next-NVFP4) | NVFP4 [`RadixArk/Qwen3.8-Flash-Next-NVFP4`](https://huggingface.co/RadixArk/Qwen3.8-Flash-Next-NVFP4) | 3, 4, 7, 11, 26 |
| GLM 5.3 Flash | [`ominn/SketchSSM-GLM-5.3-Flash-NVFP4`](https://huggingface.co/ominn/SketchSSM-GLM-5.3-Flash-NVFP4) | NVFP4 [`RedHatAI/GLM-5.3-Flash-NVFP4`](https://huggingface.co/RedHatAI/GLM-5.3-Flash-NVFP4) | 3, 4, 7, 12, 28 |
| Qwen3.5 9B | [`ominn/SketchSSM-Qwen3.5-9B-BF16`](https://huggingface.co/ominn/SketchSSM-Qwen3.5-9B-BF16) | BF16 [`Qwen/Qwen3.5-9B`](https://huggingface.co/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`](https://github.com/SNU-ARC/SketchSSM/blob/main/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](https://github.com/SNU-ARC/SketchSSM/blob/main/sketchssm/calibration/docs/new_model.md)), then run the
calibration stages:

```bash
python -m sketchssm.calibration calibrate --config my_model.yaml --out outputs/my_model
```

1. `generate`: generate continuations from WikiText-2 prompts.
2. `covariance`: replay them and collect state and query covariances.
3. `basis`: fit one shared sketch basis per head group.
4. `paired`: score each rank on validation text with output-error and gradient pairs.
5. `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:

```bash
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](https://github.com/SNU-ARC/SketchSSM/blob/main/sketchssm/calibration/README.md) 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:

```bash
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:

```bash
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](https://github.com/SNU-ARC/SketchSSM/blob/main/sketchssm/calibration/example/README.md) for
the collection recipe and the expected results of each model.

## 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 [`vllm/README.md`](https://github.com/SNU-ARC/SketchSSM/blob/main/vllm/README.md) for
kernel 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.
