PidokuInfra

Pruning, Distillation, and Architecture Optimization

Intermediate Advanced 1h Difficulty 3/5 Topic 14 of 14

Prerequisites 02, III.09


1. What is it?#

Three ways to make the model itself smaller, rather than making the same model run faster.

PRUNING       remove weights (or heads, or layers) that contribute little
DISTILLATION  train a small model to imitate a large one
ARCHITECTURE  design the model for inference efficiency from the start

These are model-modification techniques, not serving techniques. They require training compute and evaluation effort, and they change the model’s identity. But their payoff is the largest of any optimization: a model half the size is roughly twice as cheap, permanently.


2. Problem → Why → Optimization#

PROBLEM   All the serving optimizations in this section together give maybe
          5-10x. Sometimes that isn't enough.
WHY       You're still running a model that was sized for benchmark scores,
          not for your task.
OPTIMIZE  Make the model smaller in a way that preserves the capability
          YOU need.

TRADE-OFFS
  ✓ 2-10x, and it compounds with every serving optimization
  ✗ requires training compute and expertise
  ✗ requires task-specific evaluation to know what you lost
  ✗ you now own a model variant (maintenance, retraining, provenance)

WHEN TO USE   High volume, a specific task, a plateau on serving optimizations.
WHEN NOT TO   Low volume, general-purpose use, no evaluation infrastructure,
              or before you've done the serving optimizations.

3. Simple analogy#

Serving optimizations are tuning the engine. This is choosing a smaller car.

A better engine tune gets you 15% more efficiency. A car half the weight gets you 40% — but you can’t carry as much, and you had to buy a different car.

The right question is: what are you actually carrying? Most deployments use a general-purpose model for a narrow task. A 3B model fine-tuned for your task frequently beats a 70B general model on that task, at 20x the efficiency.


4. Distillation — the highest-value option#

Teacher (70B) generates outputs → student (7B) trains to match them.

Two variants:
  Response distillation:  train on the teacher's generated text (easy, common)
  Logit distillation:     train on the teacher's full output distribution
                          (more information per token, better results, needs
                          teacher logits)

Why it works better than training the small model from scratch: the teacher’s outputs are a much richer training signal than raw text. The student learns the teacher’s distribution, not just its argmax.

Typical results (task-specific):
  70B teacher → 7B student, trained on 100k teacher outputs for the task
  → student reaches 92-98% of teacher quality ON THAT TASK
  → 10x cheaper to serve
  → but much worse on everything else

That last line is the point and the caveat. Distillation trades generality for efficiency. If you serve one task, it’s an enormous win. If you serve a general assistant, it isn’t.

Practical guidance:

  • Collect real production traffic as the distillation prompt set. It’s the right distribution by construction.
  • Use the teacher’s outputs, filtered for quality.
  • Evaluate on held-out production traffic, not on public benchmarks.

5. Pruning#

UNSTRUCTURED: zero out individual weights below a threshold.
  ✓ best quality per parameter removed
  ✗ irregular sparsity gives NO speedup on GPUs (no dense-kernel benefit)
  → useful only for storage, not for inference speed

SEMI-STRUCTURED (2:4): exactly 2 of every 4 consecutive weights are zero.
  ✓ 2x tensor core throughput on Ampere+
  ✗ needs retraining/fine-tuning to recover quality
  ✗ quality cost is real (1-3% typically)
  → [EMERGING] adoption has been limited

STRUCTURED: remove entire units
  - Attention heads:   remove low-importance heads
  - FFN channels:      remove low-importance intermediate dimensions
  - LAYERS:            remove entire transformer blocks   ← the big one
  ✓ real speedup with no special kernels
  ✓ composes with everything else
  ✗ quality cost, recoverable with fine-tuning

Layer pruning is underrated. Empirically, the middle layers of large models are the most redundant:

Llama-3-70B (80 layers):
  remove 8 middle layers → 72 layers, 10% faster and smaller
  quality after brief fine-tuning: within 1-2% on most benchmarks

Published "depth pruning" results (Gromov et al., "The Unreasonable Ineffectiveness
of the Deeper Layers", 2024) show substantial layer removal with modest loss —
though performance on reasoning tasks degrades faster than on knowledge tasks.

Width pruning + distillation (NVIDIA’s Minitron approach) is the current best practice: prune, then distill from the original model to recover. Reported to reach comparable quality to training from scratch at a fraction of the compute.


6. Architecture choices that matter for inference#

If you are training or selecting a model, these decisions determine its serving cost permanently:

DECISION              INFERENCE IMPACT
GQA/MQA vs MHA        4-8x KV cache. THE biggest single lever.
MLA                   10-20x KV cache (DeepSeek). Larger still.
Sliding window        caps KV at W regardless of context.
Cross-layer KV share  2x KV. [EMERGING]
MoE                   large total params, small active params.
                      → cheap compute, expensive memory. Suits high-batch serving.
Depth vs width        deeper = more layers = more kernel launches, more KV,
                      more sequential dependency. Wider = better GPU utilization.
                      For inference, prefer wider-and-shallower at equal parameters.
Vocabulary size       bigger = fewer tokens per text (cheaper decode) but larger
                      embedding and LM head. There's an optimum.
head_dim              affects which attention kernels are available.
                      128 is well-supported; unusual values may lack fast kernels.
Multi-token predict   built-in speculative decoding.
Tied embeddings       saves V×d parameters.

“Deeper vs wider” is worth dwelling on. Two models with the same parameter count, one with 80 layers of width 8192 and one with 48 layers of width 12288:

80×8192:  more sequential steps per token, more launches, more KV (∝ L),
          worse tensor core utilization per GEMM
48×12288: fewer, larger GEMMs; less KV; fewer launches
          → typically 20-35% faster inference at equal parameters

Training dynamics favor depth; inference favors width. This is a real tension, and inference engineers should be in the room when it’s decided.


7. Performance and when each applies#

Technique              Speedup    Quality cost   Effort      Generality lost
Distillation (task)    5-20x      2-8% on task   weeks       ALL other tasks
Layer pruning + FT     1.2-1.5x   1-3%           days        some
Width pruning + FT     1.5-2.5x   2-5%           weeks       some
2:4 sparsity + FT      1.3-1.8x   1-3%           days        little
Architecture choice    1.2-8x     0 (it's a      —           depends
  (at design time)                different model)

Distillation dominates when you have a narrow task and volume. Everything else is incremental.


8. Production implications#

  • Do the serving optimizations first. They’re cheaper and reversible. Model surgery is neither.
  • Distillation requires an evaluation harness on your actual task. Without it you cannot know what you lost.
  • Use production traffic as the distillation dataset. It’s the correct distribution.
  • Model variants have ownership costs: retraining when the base updates, provenance tracking, separate evaluation, separate deployment.
  • A cascade may beat a single distilled model: route easy queries to a small model, hard ones to a large one (Section XII.09). Often gets 80% of the cost benefit with none of the quality risk on hard queries.
  • When selecting a model to deploy, weight the architecture factors. Two models with equal benchmark scores can differ 5x in serving cost.

9. Common mistakes#

Doing model surgery before serving optimization. Wrong order; much more effort.

Unstructured pruning expecting a speedup. GPUs don’t benefit from irregular sparsity.

Distilling without a task evaluation. You cannot tell what broke.

Evaluating a distilled model on public benchmarks. It was distilled for your task; measure that.

Ignoring architecture when selecting a model. KV size per token is a direct cost multiplier.

Assuming quality loss is uniform. Reasoning and long-chain tasks degrade faster than knowledge recall.


10. Hands-on exercise#

A. Layer pruning. Take a small model. Remove layers one at a time (from the middle) and measure perplexity after each removal. Plot quality vs layers removed. Where’s the cliff?

B. Distillation. Distill a 7B model into a 1B one for a narrow task (e.g. classification, or a structured extraction task). Measure quality on the task and on a general benchmark. Quantify the generality lost.

C. Architecture cost comparison. For five open models of similar benchmark quality, compute serving cost per million tokens using your Section V.15 calculator. Rank them. Does the ranking match their benchmark ranking?

D. Depth vs width. Find two models with similar parameter counts but different depth/width ratios. Measure decode throughput for each. Does the wider one win, and by how much?

E. Cascade. Build a two-model cascade: a small model answers, and a confidence signal routes uncertain queries to a large one. Measure cost and quality vs always using the large model.


11. Interview questions#

  1. When would you distill a model rather than optimize serving?
  2. Why doesn’t unstructured pruning speed up GPU inference?
  3. What is 2:4 sparsity and why hasn’t it seen wide adoption?
  4. Which layers of a large transformer are most redundant, and how do you know?
  5. Given equal parameters, would you prefer a deeper or wider model for inference? Why?
  6. Name five architecture decisions that determine a model’s serving cost.
  7. What does a model cascade buy you compared to a distilled model?

12. Further reading#

  • [ESTABLISHED] Hinton et al., “Distilling the Knowledge in a Neural Network” (2015)
  • [ESTABLISHED] Muralidharan et al., “Compact Language Models via Pruning and Knowledge Distillation” (Minitron, 2024)
  • [EMERGING] Gromov et al., “The Unreasonable Ineffectiveness of the Deeper Layers” (2024)
  • [ESTABLISHED] Frantar & Alistarh, “SparseGPT” (2023)
  • [REFERENCE] NVIDIA structured sparsity documentation
  • Next: Section VIII — Serving Systems

↑↓ navigate↵ openesc close