Serving linear-attention LLMs shifts memory pressure from ever-growing KV caches to persistent recurrent states; naive low-precision quantization of those states causes errors that accumulate across decoding steps and severely damage accuracy. The paper's core insight is that these quantization errors matter along two complementary axes—when (lifetime across steps) and where (which key rows/columns)—and that a quantization plan should allocate precision accordingly rather than apply a uniform scheme.
Key Findings
- Lifetime-aware Bit Allocation: quantization precision is assigned per state unit based on its quantization error and how long that error persists during decoding, so units that accumulate and retain larger errors receive more bits under a fixed memory budget.
- Key-Row-Aware Dual-Axis Fitting: spatial error is reduced by learning separate scales for key rows and value columns, with row scales calibrated by measured impact on output error and column scales fitted with higher weight on important rows.
- Empirical results: evaluated on Qwen-3.8B-27B and Kimi-Linear-48B-A3B-Instruct across long- and short-generation tasks, STEPQuant matches FP32-state accuracy under a nominal 6-bit budget and outperforms uniform INT8 at 4 bits; integrated kernels in SGLang yield ≈5× recurrent-state compression and up to 68.7% total serving memory reduction.
Who it helps and trade-offs
Great fit if you need to reduce GPU memory for concurrent LLM serving with linear-attention or Delta-rule recurrent states while keeping near-FP32 accuracy. It is especially useful for production inference stacks that can incorporate a calibrated quantization plan and optimized kernels. Look elsewhere if you need quantization during training, your model uses standard softmax attention (non-linear attention caches), or you cannot supply representative calibration traces—STEPQuant is a post-training, inference-focused method and requires distributional calibration and integration work for kernels.
Key method (brief)
STEPQuant computes per-unit quantization error and lifetime from calibration runs, solves a mixed-precision allocation under a bit-budget, and fits dual-axis scales (per key-row and per value-column) to minimize readout error. The approach is post-training and inference-oriented: it compresses and dequantizes recurrent states across decoding steps while limiting propagated error through lifetime- and impact-aware design.