Pruning and sparsity
Be able to prune weights and measure accuracy against speed.
Prerequisites
- FQuantisationrequired
Intuition
Pruning removes weights that contribute little. Two families with entirely different practical consequences:
| Type | What is removed | Compression | Faster? |
|---|---|---|---|
| Unstructured | individual weights, anywhere | high (50–90 %) | usually no without special support |
| Structured | whole channels, heads, layers | lower (10–50 %) | yes, directly |
| Semi-structured (2:4) | 2 of 4 adjacent weights | 50 % | yes, on Ampere+ with sparse tensor cores |
The most common disappointment: you prune 80 % of the weights away, the file becomes small — and the inference runs exactly as slowly. Unstructured sparsity requires hardware or kernel support to give speed; otherwise the zeros are multiplied just like every other number.
Formal
Magnitude pruning is the baseline: remove the weights with the smallest absolute value, layer by layer or globally. Simple and surprisingly strong.
SparseGPT and Wanda are post-training methods adapted for large language models. Wanda is strikingly simple: the importance = — the size of the weight times the norm of the corresponding activation, computed on calibration data. No retraining, no gradient computation, and it reaches 50 % sparsity with little loss.
The lottery ticket hypothesis (Frankle & Carbin 2018): in a randomly initialised network there is a small subnetwork that, trained from the same initialisation, can reach comparable performance. That explains why pruning works — but it does not give a practical route to faster training, since the subnetwork is found by first training the whole network.
The order of decisions in practice: quantisation first (simpler, nearly always gives speed), then structured pruning if more is needed, and unstructured only if the runtime has support for it. Combining quantisation and sparsity works but requires careful measurement — the losses do not add linearly.
Code
import torch
def magnitude_prune(model, share=0.5, global_=True):
"""Zero the smallest weights out. Returns the actual sparsity."""
weights = [p for n, p in model.named_parameters() if p.dim() > 1]
if global_:
every = torch.cat([p.detach().abs().flatten() for p in weights])
threshold = torch.quantile(every, share)
for p in weights:
p.data[p.abs() < threshold] = 0
else:
for p in weights:
threshold = torch.quantile(p.detach().abs().flatten(), share)
p.data[p.abs() < threshold] = 0
zeros = sum(int((p == 0).sum()) for p in weights)
total = sum(p.numel() for p in weights)
return zeros / total
def wanda_score(W, activation_norm):
"""Wanda: |w| * ||x||. activation_norm: (in_features,) from calibration data."""
return W.abs() * activation_norm.unsqueeze(0)
# ALWAYS measure three things after pruning:
# 1. the sparsity (the share of zeros) 2. the quality on your eval 3. the actual tokens/s
# If (3) is unchanged you have only made the model worse.
Mastery means
- Prunes weights in a structured and an unstructured way
- Measures the accuracy-against-speed trade-off
- Knows when sparsity actually gives faster inference
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — A Simple and Effective Pruning Approach for Large Language Models (Wanda) — arXiv (open access; licence per article)
- arXiv — SparseGPT: Massive Language Models Can Be Accurately Pruned in One-Shot — arXiv (open access; licence per article)
- arXiv — The Lottery Ticket Hypothesis — arXiv (open access; licence per article)