Sparsity & Pruning advanced 8 min read 6 flashcards

Contextual Sparsity and Activation Predictors

If you knew in advance which neurons a token will activate you could avoid loading the rest, so the whole technique reduces to building a predictor that is cheaper than the work it saves and almost never wrong.

The awkward part of activation sparsity is the ordering. To know that neuron \(j\) is inactive you compute \(\sigma(W_{\text{gate}}[j,:]\,x)\), which means reading the row you were hoping to skip. Knowing the answer after paying for it is worthless.

Contextual sparsity is the escape. The claim is that for a given input there exists a small, input-dependent set of attention heads and MLP neurons whose output approximates the dense layer, and that this set is predictable from the layer's input by something far cheaper than the layer itself (Liu et al., 2023, Deja Vu: Contextual Sparsity for Efficient LLMs at Inference Time, PMLR 202:21631-21657, arXiv:2310.17157). On OPT-175B, over 80 percent of attention heads can be silenced for a given token and over 95 percent of MLP parameters zeroed, for roughly 85 percent total sparsity, with zero-shot average accuracy flat until about 75 percent.

The predictor and its budget

Deja Vu's predictor is a small two-layer network per layer, trained offline on recorded activations to output a score per neuron; you take the top-\(k\) scores as the predicted active set. Its economics are a straight inequality. Let \(B_{\text{dense}}\) be the bytes the layer would read, \(s\) the sparsity it predicts, and \(B_{\text{pred}}\) the bytes the predictor itself reads. The technique wins only when

\[B_{\text{pred}} + (1-s)\,B_{\text{dense}} < B_{\text{dense}}\]

so the predictor's budget is \(s \cdot B_{\text{dense}}\) and nothing more. For an 11 GB feedforward stack at 90 percent predicted sparsity that is roughly 10 GB of headroom, which is why a few hundred megabytes of predictor is affordable and a predictor the size of a transformer layer is not.

The second trick is latency hiding. A strictly sequential pipeline of predict-then-gather-then-compute adds the predictor's latency to the critical path at every layer. Deja Vu instead predicts the sparsity of layer \(\ell+1\) from the input to layer \(\ell\) and runs the prediction asynchronously, explicitly modelled on a hardware branch predictor. That works because the residual stream changes slowly between adjacent blocks: \(x_{\ell+1} = x_\ell + f_\ell(x_\ell)\) with an update that is small relative to the accumulated stream, so the input to layer \(\ell\) is a good proxy for the input to layer \(\ell+1\). The same residual-norm argument underwrites depth pruning, from the other direction.

Static locality: hot and cold neurons

A predictor is one way to exploit input dependence. The other is to notice that the dependence is not uniform. Neuron activation across inputs follows a power law: a small subset fires for almost every token, and the long tail fires only for particular inputs. PowerInfer builds its whole design on that split, preloading the hot neurons into GPU memory permanently and computing the cold ones on the CPU, which lets a single consumer GPU serve models far larger than its VRAM and beats llama.cpp by up to 11.69x on an RTX 4090 (Song et al., 2024, PowerInfer: Fast Large Language Model Serving with a Consumer-grade GPU, SOSP 2024, arXiv:2312.12456).

The two approaches compose: the static split decides where a neuron lives, the predictor decides whether this token needs it.

Precision, recall, and which error hurts

A predictor makes two kinds of mistake, and they are not symmetric.

A false positive, marking an inactive neuron active, costs bandwidth. You read a row you did not need. Quality is untouched.

A false negative, marking an active neuron inactive, drops a real contribution from the output. One dropped neuron among 14,336 is noise. A systematic bias, say the predictor reliably misses the neurons that fire on rare tokens or on code, is a quality regression that no aggregate perplexity check will surface.

So predictors are tuned for recall, deliberately over-predicting the active set. A predictor at 95 percent recall and 60 percent precision is usually the better deployment than the reverse, because the first wastes bandwidth and the second silently damages outputs.

When it breaks

Batching dissolves it. The active sets of different tokens in a batch are different, and the layer must read the union. Beyond a modest batch size the union approaches the full matrix, MLP sparsity vanishes, and the gather machinery becomes pure overhead (Shrestha et al., 2025, Polar Sparsity: High Throughput Batched LLM Inferencing with Scalable Contextual Sparsity, arXiv:2505.14884).

The predictor is a trained model with a training distribution. Predictors fitted on recorded activations from web text carry that distribution's idea of which neurons matter. Tool-call formatting, a low-resource language or an unusual code dialect can shift the active set in ways the predictor has never scored.

Per-layer predictors are per-checkpoint artefacts. Fine-tune the model and the activation statistics move, so the predictors must be refitted. That is a real operational cost in any pipeline that ships LoRA adapters or continually post-trains.

The asynchronous trick assumes a slowly changing stream. Early layers, where the residual norm is still small relative to each block's update, are exactly where the look-ahead approximation is weakest.

References and further reading

Every source this page cites, in the order it cites them. All of them open in a new tab.

  1. Liu et al., 2023, Deja Vu: Contextual Sparsity for Efficient LLMs at Inference Time, PMLR 202:21631-21657, arXiv:2310.17157 arxiv.org
  2. Song et al., 2024, PowerInfer: Fast Large Language Model Serving with a Consumer-grade GPU, SOSP 2024, arXiv:2312.12456 arxiv.org
  3. Shrestha et al., 2025, Polar Sparsity: High Throughput Batched LLM Inferencing with Scalable Contextual Sparsity, arXiv:2505.14884 arxiv.org
Check yourself

6 flashcards for this concept

Click a card to reveal the answer.

Drill the whole track