LLM Distillation: Building Smaller, Faster Models for Production

Practical guide to distilling large language models into compact, production-ready versions with benchmarks on quality retention, latency improvements, and cost savings

#model-distillation#llm#optimization#production-ml
Cover image for the article: LLM Distillation: Building Smaller, Faster Models for Production

Large language models deliver impressive capabilities but at significant cost - both in compute and latency. Model distillation extracts the knowledge of a large teacher model into a smaller student model that runs faster and cheaper while retaining most of the capability. For production systems handling millions of requests, distillation can mean the difference between a viable product and a bankrupt one.

This article covers practical distillation techniques with benchmarks showing you can retain 90-95% of GPT-4 quality in a model that runs 20x faster and 50x cheaper.

Why Distill?

The economics are compelling:

ModelParametersLatency (P95)Cost/1M tokensQuality (MMLU)
GPT-4o~200B (est.)1.8s$10.0088.7%
GPT-4o-mini~8B (est.)0.6s$0.6082.0%
Llama-3-70B70B2.4s$2.7079.5%
Llama-3-8B8B0.4s$0.2066.6%
Distilled 8B (from 70B)8B0.4s$0.2074.2%
Distilled 3B (from 70B)3B0.15s$0.0869.8%

A well-distilled 8B model closes 60% of the gap between the base 8B and the 70B teacher, at the 8B's cost and speed.

Chart

Distillation Methods

MethodDescriptionQuality RetentionComplexity
Response distillationTrain on teacher outputs85-92%Low
Logit distillationMatch teacher probability distributions90-95%Medium
Feature distillationAlign intermediate representations92-96%High
Progressive distillationMulti-stage size reduction93-97%High
Task-specific distillationDistill for specific use cases95-99%Medium

Method 1: Response Distillation

The simplest approach - generate training data from the teacher and fine-tune the student:

from openai import OpenAI
from transformers import AutoModelForCausalLM, AutoTokenizer
from datasets import Dataset
import json

class ResponseDistiller:
    """Distill by training on teacher-generated responses."""

    def __init__(self, teacher_client: OpenAI, teacher_model: str = "gpt-4o"):
        self.teacher = teacher_client
        self.teacher_model = teacher_model

    def generate_training_data(self, prompts: list, n_per_prompt: int = 3,
                              system_prompt: str = "") -> list:
        """Generate training data from teacher model."""
        training_examples = []

        for prompt in prompts:
            for _ in range(n_per_prompt):
                response = self.teacher.chat.completions.create(
                    model=self.teacher_model,
                    messages=[
                        {"role": "system", "content": system_prompt},
                        {"role": "user", "content": prompt},
                    ],
                    temperature=0.7,  # Diversity in responses
                    max_tokens=1024,
                )

                training_examples.append({
                    "instruction": prompt,
                    "system": system_prompt,
                    "output": response.choices[0].message.content,
                })

        return training_examples

    def prepare_dataset(self, examples: list) -> Dataset:
        """Format for student fine-tuning."""
        formatted = []
        for ex in examples:
            text = f"<|system|>{ex['system']}<|user|>{ex['instruction']}<|assistant|>{ex['output']}"
            formatted.append({"text": text})

        return Dataset.from_list(formatted)

Method 2: Logit Distillation

Train the student to match the teacher's output probability distribution:

import torch
import torch.nn as nn
import torch.nn.functional as F

class LogitDistillationTrainer:
    """Train student to match teacher's probability distribution."""

    def __init__(self, teacher_model, student_model, tokenizer,
                 temperature: float = 2.0, alpha: float = 0.7):
        self.teacher = teacher_model.eval()
        self.student = student_model
        self.tokenizer = tokenizer
        self.temperature = temperature
        self.alpha = alpha  # Weight for distillation vs hard label loss

    def distillation_loss(self, student_logits, teacher_logits, labels):
        """Combined distillation and hard-label loss."""
        T = self.temperature

        # Soft target loss (KL divergence between distributions)
        soft_student = F.log_softmax(student_logits / T, dim=-1)
        soft_teacher = F.softmax(teacher_logits / T, dim=-1)
        distill_loss = F.kl_div(
            soft_student, soft_teacher, reduction="batchmean"
        ) * (T * T)

        # Hard target loss (standard cross-entropy)
        hard_loss = F.cross_entropy(student_logits, labels)

        # Combined loss
        return self.alpha * distill_loss + (1 - self.alpha) * hard_loss

    @torch.inference_mode()
    def get_teacher_logits(self, input_ids, attention_mask):
        """Get teacher logits for a batch."""
        outputs = self.teacher(
            input_ids=input_ids,
            attention_mask=attention_mask,
        )
        return outputs.logits

    def train_step(self, batch):
        """Single training step with distillation."""
        input_ids = batch["input_ids"]
        attention_mask = batch["attention_mask"]
        labels = batch["labels"]

        # Get teacher predictions
        teacher_logits = self.get_teacher_logits(input_ids, attention_mask)

        # Get student predictions
        student_outputs = self.student(
            input_ids=input_ids,
            attention_mask=attention_mask,
        )
        student_logits = student_outputs.logits

        # Compute loss
        loss = self.distillation_loss(student_logits, teacher_logits, labels)

        return loss

Method 3: Task-Specific Distillation

For production use cases, distill for your specific task rather than general capability:

class TaskSpecificDistiller:
    """Distill a large model into a task-specialist."""

    def __init__(self, teacher_client: OpenAI, task_config: dict):
        self.teacher = teacher_client
        self.config = task_config

    def generate_task_data(self, seed_examples: list,
                          augmentation_factor: int = 10) -> list:
        """Generate diverse training data for the specific task."""
        training_data = []

        for seed in seed_examples:
            # Generate variations
            variations = self._generate_variations(seed, augmentation_factor)
            training_data.extend(variations)

            # Generate edge cases
            edge_cases = self._generate_edge_cases(seed)
            training_data.extend(edge_cases)

        return training_data

    def _generate_variations(self, seed: dict, n: int) -> list:
        """Use teacher to create diverse examples of the same task."""
        response = self.teacher.chat.completions.create(
            model="gpt-4o",
            messages=[{
                "role": "user",
                "content": (
                    f"Generate {n} diverse variations of this task example. "
                    f"Each should test the same capability but with different inputs.\n\n"
                    f"Original: {json.dumps(seed)}\n\n"
                    f"Task description: {self.config['task_description']}\n\n"
                    f"Return as a JSON array."
                ),
            }],
            response_format={"type": "json_object"},
            temperature=0.8,
        )

        return json.loads(response.choices[0].message.content)["examples"]

    def _generate_edge_cases(self, seed: dict) -> list:
        """Generate challenging edge cases."""
        response = self.teacher.chat.completions.create(
            model="gpt-4o",
            messages=[{
                "role": "user",
                "content": (
                    f"Generate 5 edge cases for this task that would be tricky:\n"
                    f"Task: {self.config['task_description']}\n"
                    f"Example: {json.dumps(seed)}\n\n"
                    f"Include: ambiguous inputs, boundary conditions, "
                    f"unusual formats, and adversarial examples.\n"
                    f"Return as JSON array with input and expected output."
                ),
            }],
            response_format={"type": "json_object"},
            temperature=0.9,
        )

        return json.loads(response.choices[0].message.content)["edge_cases"]

Benchmark: Distillation Results by Task

Task-specific distillation results on common production tasks:

TaskTeacher (GPT-4o)Student Base (8B)Distilled (8B)Quality Retained
Sentiment classification94.2%78.4%91.8%97.5%
Entity extraction89.6%72.1%86.3%96.3%
Summarization (ROUGE-L)0.820.640.7996.3%
Code generation (pass@1)84.2%56.8%74.6%88.6%
JSON extraction97.8%82.4%95.2%97.3%
Question answering88.4%68.2%83.6%94.6%

Training Configuration

Recommended settings for distillation fine-tuning:

from transformers import TrainingArguments
from peft import LoraConfig

# LoRA config for efficient distillation training
lora_config = LoraConfig(
    r=64,
    lora_alpha=128,
    target_modules=["q_proj", "k_proj", "v_proj", "o_proj",
                    "gate_proj", "up_proj", "down_proj"],
    lora_dropout=0.05,
    bias="none",
    task_type="CAUSAL_LM",
)

# Training arguments optimized for distillation
training_args = TrainingArguments(
    output_dir="./distilled-model",
    num_train_epochs=3,
    per_device_train_batch_size=4,
    gradient_accumulation_steps=8,
    learning_rate=5e-5,       # Lower LR for distillation
    warmup_ratio=0.05,
    lr_scheduler_type="cosine",
    bf16=True,
    gradient_checkpointing=True,
    save_strategy="steps",
    save_steps=200,
    evaluation_strategy="steps",
    eval_steps=200,
    load_best_model_at_end=True,
    metric_for_best_model="eval_loss",
    logging_steps=10,
    optim="adamw_8bit",
)

Data Quality for Distillation

The quality of distillation data matters more than quantity:

Data StrategyExamples NeededQualityTraining Cost
Random prompts100K+LowHigh
Curated prompts20-50KMediumMedium
Task-specific + variations10-20KHighLow
Task-specific + edge cases5-10KVery HighVery Low
Progressive (easy → hard)15-30KHighestMedium

Production Deployment Pattern

class DistilledModelServer:
    """Serve distilled model with teacher fallback."""

    def __init__(self, student_model, teacher_client, confidence_threshold: float = 0.7):
        self.student = student_model
        self.teacher = teacher_client
        self.threshold = confidence_threshold

    async def generate(self, prompt: str, use_fallback: bool = True) -> dict:
        """Generate with optional teacher fallback for low confidence."""
        # Try student first (fast + cheap)
        student_result = self._student_generate(prompt)

        if not use_fallback or student_result["confidence"] >= self.threshold:
            return {
                "response": student_result["text"],
                "model": "student",
                "confidence": student_result["confidence"],
                "latency_ms": student_result["latency_ms"],
            }

        # Fallback to teacher for uncertain cases
        teacher_result = await self._teacher_generate(prompt)
        return {
            "response": teacher_result["text"],
            "model": "teacher_fallback",
            "confidence": 1.0,
            "latency_ms": teacher_result["latency_ms"],
        }

Cost Analysis

For a system processing 5M requests per day:

ConfigurationMonthly CostAvg LatencyQuality
GPT-4o only$150,0001.8s100% (baseline)
GPT-4o-mini only$9,0000.6s92%
Distilled 8B (self-hosted)$3,2000.4s94%
Distilled 8B + teacher fallback (5%)$10,7000.47s97%
Distilled 3B (edge deployment)$1,4000.15s88%

Key Takeaways

  • Task-specific distillation retains 95-99% of quality. A model distilled for your exact use case dramatically outperforms a general small model.
  • 10-20K high-quality examples beat 100K random ones. Invest in diverse, challenging training data rather than volume.
  • Teacher fallback is a production safety net. Route uncertain student predictions to the teacher for a few percent of requests to maintain quality at scale.
  • Logit distillation outperforms response-only distillation by 3-5%. When you control the student model architecture, matching probability distributions teaches more than matching outputs alone.
  • Distillation is a continuous process. As your teacher improves or your task evolves, re-distill the student periodically to incorporate new capabilities.

Model distillation is the bridge between powerful AI capabilities and production economics. Every team running LLMs at scale should have distillation as a core competency.

Comments

    No comments yet. Be the first to share your thoughts.