Knowledge Distillation

orchestra-research/ai-research-skills/19-emerging-techniques/knowledge-distillation

by orchestra-research773a52944ba4MIT13K starsListed Oct 8, 2026Updated Oct 8, 2026Repository updated 3 months ago

Compress large language models using knowledge distillation from teacher to student models. Use when deploying smaller models with retained performance, transferring GPT-4 capabilities to open-source models, or reducing inference costs. Covers temperature scaling, soft targets, reverse KLD, logit distillation, and MiniLLM training strategies.

Instructions onlyAI & Agents
AI-generated overview

Guides knowledge distillation of large language models into smaller student models using temperature scaling, soft targets and reverse KLD.

What it does
This skill provides instructions and code patterns for compressing large language models into smaller student models through knowledge distillation. It covers temperature scaling, soft targets, forward versus reverse KLD (MiniLLM), logit distillation, response distillation, multi-teacher distillation and two-stage training, with a complete Trainer-based training script. It also gives guidance on hyperparameters, teacher-to-student size ratios, data quality and evaluation of the distilled model.
When to use it
Use it when deploying smaller models that should retain most of a larger teacher's performance, transferring capabilities from proprietary models to open-source ones, or reducing inference costs. It also fits creating specialized or improved small models via distillation.
Requirements
Requires Python with transformers, torch and datasets (plus accelerate, deepspeed and wandb for training), and optionally the MiniLLM implementation from the LMOps repository. Access to teacher and student model weights is needed. It ships no scripts; it is instructions and code examples only.

Knowledge Distillation: Compressing LLMs

When to Use This Skill

Use Knowledge Distillation when you need to:

  • Compress models from 70B → 7B while retaining 90%+ performance
  • Transfer capabilities from proprietary models (GPT-4) to open-source (LLaMA, Mistral)
  • Reduce inference costs by deploying smaller student models
  • Create specialized models by distilling domain-specific knowledge
  • Improve small models using synthetic data from large teachers

Key Techniques: Temperature scaling, soft targets, reverse KLD (MiniLLM), logit distillation, response distillation

Papers: Hinton et al. 2015 (arXiv 1503.02531), MiniLLM (arXiv 2306.08543), KD Survey (arXiv 2402.13116)

Installation

bash
# Standard transformerspip install transformers datasets accelerate
# For trainingpip install torch deepspeed wandb
# Optional: MiniLLM implementationgit clone https://github.com/microsoft/LMOpscd LMOps/minillmpip install -e .

Quick Start

Basic Knowledge Distillation

python
import torchimport torch.nn.functional as Ffrom transformers import AutoModelForCausalLM, AutoTokenizer, Trainer, TrainingArguments
# 1. Load teacher (large) and student (small) modelsteacher = AutoModelForCausalLM.from_pretrained(    "meta-llama/Llama-2-70b-hf",  # Large teacher    torch_dtype=torch.float16,    device_map="auto")
student = AutoModelForCausalLM.from_pretrained(    "meta-llama/Llama-2-7b-hf",  # Small student    torch_dtype=torch.float16,    device_map="cuda:0")
tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-2-70b-hf")
# 2. Define distillation lossdef distillation_loss(student_logits, teacher_logits, labels, temperature=2.0, alpha=0.5):    """    Combine hard loss (cross-entropy) with soft loss (KL divergence).
    Args:        temperature: Softens probability distributions (higher = softer)        alpha: Weight for distillation loss (1-alpha for hard loss)    """    # Hard loss: Standard cross-entropy with true labels    hard_loss = F.cross_entropy(student_logits.view(-1, student_logits.size(-1)), labels.view(-1))
    # Soft loss: KL divergence between student and teacher    soft_targets = F.softmax(teacher_logits / temperature, dim=-1)    soft_student = F.log_softmax(student_logits / temperature, dim=-1)    soft_loss = F.kl_div(soft_student, soft_targets, reduction='batchmean') * (temperature ** 2)
    # Combined loss    return alpha * soft_loss + (1 - alpha) * hard_loss
# 3. Training loopfor batch in dataloader:    # Teacher forward (no grad)    with torch.no_grad():        teacher_outputs = teacher(**batch)        teacher_logits = teacher_outputs.logits
    # Student forward    student_outputs = student(**batch)    student_logits = student_outputs.logits
    # Compute distillation loss    loss = distillation_loss(        student_logits,        teacher_logits,        batch['labels'],        temperature=2.0,        alpha=0.7  # 70% soft, 30% hard    )
    # Backward and optimize    loss.backward()    optimizer.step()    optimizer.zero_grad()

MiniLLM (Reverse KLD)

Source: arXiv 2306.08543 (2024)

Innovation: Use reverse KLD instead of forward KLD for better generative model distillation.

python
def reverse_kl_loss(student_logits, teacher_logits, temperature=1.0):    """    Reverse KL divergence: KL(Teacher || Student)    Better for generative models than forward KL.    """    # Teacher distribution (target)    p_teacher = F.softmax(teacher_logits / temperature, dim=-1)
    # Student distribution (model)    log_p_student = F.log_softmax(student_logits / temperature, dim=-1)
    # Reverse KL: Sum over teacher, student learns to cover teacher's modes    reverse_kl = -(p_teacher * log_p_student).sum(dim=-1).mean()
    return reverse_kl * (temperature ** 2)
# Training with MiniLLMfor batch in dataloader:    with torch.no_grad():        teacher_logits = teacher(**batch).logits
    student_logits = student(**batch).logits
    # Reverse KLD (better for generation)    loss = reverse_kl_loss(student_logits, teacher_logits, temperature=1.0)
    loss.backward()    optimizer.step()

Why reverse KL?

  • Forward KL (standard): Student learns to match teacher's mean
  • Reverse KL (MiniLLM): Student learns to cover all teacher's modes
  • Better for diverse text generation

Response Distillation

python
# Generate synthetic data from teacher, train student to imitate
# 1. Generate synthetic responses from teacherprompts = ["Explain AI:", "What is ML?", "Define NLP:"]
teacher_responses = []for prompt in prompts:    inputs = tokenizer(prompt, return_tensors='pt').to(teacher.device)    outputs = teacher.generate(**inputs, max_new_tokens=256, do_sample=True, temperature=0.7)    response = tokenizer.decode(outputs[0], skip_special_tokens=True)    teacher_responses.append(response)
# 2. Train student on teacher's responses (standard fine-tuning)train_dataset = [    {"text": f"{prompt}\n{response}"}    for prompt, response in zip(prompts, teacher_responses)]
# 3. Fine-tune studenttrainer = Trainer(    model=student,    args=TrainingArguments(output_dir="./student", num_train_epochs=3, learning_rate=2e-5),    train_dataset=train_dataset,)trainer.train()

Core Concepts

1. Temperature Scaling

Purpose: Soften probability distributions to expose teacher's uncertainty.

python
# Low temperature (T=1): Sharp distributionlogits = [3.0, 2.0, 1.0]probs_T1 = softmax(logits / 1.0)  # [0.67, 0.24, 0.09]
# High temperature (T=4): Soft distributionprobs_T4 = softmax(logits / 4.0)  # [0.42, 0.34, 0.24]
# Higher T reveals more information about relative rankings

Rule: Use T=2-5 for distillation (2 is common default).

2. Loss Function Components

python
# Total loss = alpha * soft_loss + (1 - alpha) * hard_loss
# Soft loss: Learn from teacher's knowledgesoft_loss = KL(student || teacher)
# Hard loss: Learn from ground truth labelshard_loss = CrossEntropy(student_output, true_labels)
# Typical values:alpha = 0.5  # Balancedalpha = 0.7  # More emphasis on teacheralpha = 0.3  # More emphasis on labels

3. Forward vs Reverse KLD

python
# Forward KL: KL(Student || Teacher)# - Student matches teacher's average behavior# - Mode-seeking: Student focuses on teacher's highest probability modes# - Good for classification
# Reverse KL: KL(Teacher || Student)# - Student covers all of teacher's behaviors# - Mode-covering: Student learns diverse behaviors# - Good for generation (MiniLLM)

Training Strategies

Strategy 1: Logit Distillation

python
# Train student to match teacher's logits directly
def logit_distillation_trainer(student, teacher, dataloader, temperature=2.0):    optimizer = torch.optim.AdamW(student.parameters(), lr=2e-5)
    for epoch in range(3):        for batch in dataloader:            # Get logits            with torch.no_grad():                teacher_logits = teacher(**batch).logits
            student_logits = student(**batch).logits
            # MSE on logits (alternative to KLD)            loss = F.mse_loss(student_logits, teacher_logits)
            # Or use KLD            # loss = F.kl_div(            #     F.log_softmax(student_logits/temperature, dim=-1),            #     F.softmax(teacher_logits/temperature, dim=-1),            #     reduction='batchmean'            # ) * (temperature ** 2)
            loss.backward()            optimizer.step()            optimizer.zero_grad()
    return student

Strategy 2: Two-Stage Distillation

python
# Stage 1: Distill from teacherstudent = distill(teacher, student, epochs=5)
# Stage 2: Fine-tune on task-specific datastudent = fine_tune(student, task_data, epochs=3)
# Results in better task performance than single-stage

Strategy 3: Multi-Teacher Distillation

python
# Learn from multiple expert teachers
def multi_teacher_distillation(student, teachers, batch):    """Distill from ensemble of teachers."""    teacher_logits_list = []
    # Get logits from all teachers    with torch.no_grad():        for teacher in teachers:            logits = teacher(**batch).logits            teacher_logits_list.append(logits)
    # Average teacher predictions    avg_teacher_logits = torch.stack(teacher_logits_list).mean(dim=0)
    # Student learns from ensemble    student_logits = student(**batch).logits    loss = F.kl_div(        F.log_softmax(student_logits, dim=-1),        F.softmax(avg_teacher_logits, dim=-1),        reduction='batchmean'    )
    return loss

Production Deployment

Complete Training Script

python
from transformers import Trainer, TrainingArguments, DataCollatorForLanguageModeling
def train_distilled_model(    teacher_name="meta-llama/Llama-2-70b-hf",    student_name="meta-llama/Llama-2-7b-hf",    output_dir="./distilled-llama-7b",    temperature=2.0,    alpha=0.7,):    # Load models    teacher = AutoModelForCausalLM.from_pretrained(teacher_name, torch_dtype=torch.float16, device_map="auto")    student = AutoModelForCausalLM.from_pretrained(student_name, torch_dtype=torch.float16)    tokenizer = AutoTokenizer.from_pretrained(teacher_name)
    # Custom trainer with distillation    class DistillationTrainer(Trainer):        def compute_loss(self, model, inputs, return_outputs=False):            # Student forward            outputs_student = model(**inputs)            student_logits = outputs_student.logits
            # Teacher forward (no grad)            with torch.no_grad():                outputs_teacher = teacher(**inputs)                teacher_logits = outputs_teacher.logits
            # Distillation loss            soft_targets = F.softmax(teacher_logits / temperature, dim=-1)            soft_student = F.log_softmax(student_logits / temperature, dim=-1)            soft_loss = F.kl_div(soft_student, soft_targets, reduction='batchmean') * (temperature ** 2)
            # Hard loss            hard_loss = outputs_student.loss
            # Combined            loss = alpha * soft_loss + (1 - alpha) * hard_loss
            return (loss, outputs_student) if return_outputs else loss
    # Training arguments    training_args = TrainingArguments(        output_dir=output_dir,        num_train_epochs=3,        per_device_train_batch_size=4,        gradient_accumulation_steps=8,        learning_rate=2e-5,        warmup_steps=500,        logging_steps=100,        save_steps=1000,        bf16=True,        gradient_checkpointing=True,    )
    # Train    trainer = DistillationTrainer(        model=student,        args=training_args,        train_dataset=train_dataset,        data_collator=DataCollatorForLanguageModeling(tokenizer, mlm=False),    )
    trainer.train()    student.save_pretrained(output_dir)    tokenizer.save_pretrained(output_dir)
# Usagetrain_distilled_model(    teacher_name="meta-llama/Llama-2-70b-hf",    student_name="meta-llama/Llama-2-7b-hf",    temperature=2.0,    alpha=0.7)

Best Practices

1. Hyperparameter Selection

python
# TemperatureT = 1.0  # Sharp (less knowledge transfer)T = 2.0  # Standard (good balance)T = 5.0  # Soft (more knowledge transfer)
# Alpha (weight)alpha = 0.5  # Balancedalpha = 0.7  # Emphasize teacher knowledgealpha = 0.9  # Strong distillation
# Rule: Higher T + higher alpha = stronger distillation

2. Model Size Ratio

python
# Good ratios (teacher/student)70B / 7B = 10×    # Excellent13B / 1B = 13×    # Good7B / 1B = 7×      # Acceptable
# Avoid too large gap70B / 1B = 70×    # Too large, ineffective

3. Data Quality

python
# Best: Use teacher-generated data + real datatrain_data = {    "teacher_generated": 70%,  # Diverse, high-quality    "real_data": 30%            # Ground truth}
# Avoid: Only real data (doesn't utilize teacher fully)

Evaluation

python
from transformers import pipeline
# Compare student vs teacherteacher_pipe = pipeline("text-generation", model=teacher)student_pipe = pipeline("text-generation", model=student)
prompts = ["Explain quantum computing:", "What is AI?"]
for prompt in prompts:    teacher_out = teacher_pipe(prompt, max_new_tokens=100)    student_out = student_pipe(prompt, max_new_tokens=100)
    print(f"Prompt: {prompt}")    print(f"Teacher: {teacher_out[0]['generated_text']}")    print(f"Student: {student_out[0]['generated_text']}")    print(f"Match quality: {calculate_similarity(teacher_out, student_out):.2f}")

Resources

Source and attribution

Source:orchestra-research/ai-research-skillsin19-emerging-techniques/knowledge-distillationat commit773a529

License: MIT

Content belongs to its original authors. SourceWeft indexes it from a public repository.

Report or request removal