🤖 AI Summary
This study addresses the weight transfer bottleneck caused by GPU memory insufficiency during LLM decoding and the accuracy degradation inherent in conventional activation sparsification. To overcome these limitations, we propose SpAx, a method that transcends binary sparsity trade-offs by introducing a tri-state mechanism—full retention, approximate compression, and direct omission. SpAx hierarchically optimizes weight retrieval based on activation magnitudes, exploiting the error-attenuation characteristics of low-magnitude activations to achieve an optimal speed–accuracy trade-off. When integrated with CPU/Flash storage offloading and quantized inference, SpAx yields an average 3.86× speedup (up to 5.57×) under CPU offloading while constraining perplexity degradation to below 10%, significantly reducing latency while effectively preserving model quality.
📝 Abstract
Deploying LLMs on consumer-grade GPUs with insufficient memory to hold their weights can result in prohibitively slow inference, because decoding repeatedly transfers offloaded weights from system RAM or flash storage into GPU at much lower bandwidth than local GPU-memory access. Activation sparsity reduces these transfers by skipping weights associated with zero or near-zero activations. However, as more activation contributions are omitted, model quality eventually degrades rapidly, indicating that weights associated with small-magnitude activations collectively influence model quality sharply. In this work, we improve the trade-off between model quality and decoding performance when exploiting activation sparsity. Our key idea is to replace the binary choice of whether or not to read a weight with three options: fully retain it, approximate it using a compressed weight representation, or omit it entirely. SpAx skips weights associated with activations closest to zero, reads approximate weights for smaller-magnitude activations, and reads original weights for the largest-magnitude activations. 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 speeds up decoding by 3.86X on average (up to 5.57X) with 16-bit weights and 2.06X (up to 2.74X) with 4-bit weights, at a WikiText-2 perplexity increase of at most 10%. With weights offloaded to flash storage, the speedups are 3.31X on average (up to 4.81X) and 1.54X (up to 2.03X).