SpAx triples weight-read options to speed LLM decoding on offloaded weights by up to 5.57x with at most 10% WikiText-2 perplexity increase
Related research and updatesSynopsis
The work introduces SpAx, which replaces the binary choice of whether to read a weight under activation sparsity with three options—fully skipping weights whose activations are closest to zero, reading compressed approximate weights for smaller-magnitude activations, and reading original weights for the largest-magnitude activations—thereby speeding up decoding when weights are offloaded: with weights offloaded to CPU memory, 3.86x average (up to 5.57x) for 16-bit weights and 2.06x average (up to 2.74x) for 4-bit weights; with weights offloaded to flash storage, 3.31x average (up to 4.81x) and 1.54x average (up to 2.03x), at a WikiText-2 perplexity increase of at most 10%.
Figure 1: Llama-3.1-8B (BF16) with all weights offloaded, at different sparsity levels. (a), (b): Time per decoded token with the weights offloaded to CPU memory and to flash storage. (c): Perplexity increase relative to dense on the WikiText-2 test set. (See Section 6 for the experiment setup.)
arXivInterpretation
SpAx expands the weight-read decision in activation sparsity from binary to three tiers: fully retain, approximate with a compressed weight representation, or omit entirely. Prior activation-sparsity methods chose only between reading and skipping a weight; this work adds an intermediate tier so that weights tied to smaller-magnitude activations participate in computation in compressed form. The abstract states the three-tier strategy explicitly and gives its rationale: smaller-magnitude activations attenuate the errors introduced by approximate weights, while compressed weight representations require fewer bytes to be transferred.
With weights offloaded to CPU memory, SpAx reports 3.86x average speedup (up to 5.57x) for 16-bit weights and 2.06x average (up to 2.74x) for 4-bit weights. The result targets the setting where consumer-grade GPU memory cannot hold model weights and decoding repeatedly transfers weights from system RAM, directly improving decoding performance in that constrained deployment. The abstract reports both average and peak figures together with the corresponding weight-bit-width conditions.
With weights offloaded to flash storage, SpAx reports 3.31x average speedup (up to 4.81x) for 16-bit weights and 1.54x average (up to 2.03x) for 4-bit weights. Flash storage offers lower bandwidth than system RAM, so this result indicates the three-tier strategy also yields speedups over a narrower transfer channel. The abstract lists average and peak speedups separately for the flash-storage scenario at 16-bit and 4-bit weights.
The quality cost accompanying these speedups is quantified as a WikiText-2 perplexity increase of at most 10%. The work reports speed and quality within a single trade-off framing rather than presenting speedup numbers alone. The abstract uses WikiText-2 perplexity as the quality metric and states an upper bound of 'at most 10%'.
Perspective
The work targets decoding scenarios where weights cannot all reside in GPU memory and must be repeatedly transferred from system RAM or flash storage, applicable to consumer-grade GPU deployment. Its benefit is bounded by a WikiText-2 perplexity increase of at most 10%, and speedups are reported separately for 16-bit and 4-bit weights and for CPU-memory versus flash-storage offload targets.
The abstract does not state the model sizes or families tested, the specific baseline configurations, the form and compression ratio of the compressed weight representation, or how the perplexity increase distributes across 16-bit and 4-bit weights. WikiText-2 perplexity is the only reported quality metric, so whether the conclusion extends to downstream tasks remains open. The abstract also does not report what fraction of end-to-end latency is attributable to non-weight-transfer components.
