student teacher

**Student-teacher learning** trains a **smaller student model to mimic a larger teacher model's behavior** — enabling deployment of compact, efficient models that retain much of the teacher's capability through knowledge distillation, intermediate layer matching, and response imitation. **What Is Student-Teacher Learning?** - **Definition**: Transfer knowledge from large (teacher) to small (student). - **Goal**: Smaller model with similar performance. - **Methods**: Logit matching, feature distillation, response copying. - **Applications**: Compression, deployment, efficient inference. **Why Student-Teacher** - **Deployment**: Large models too expensive for production. - **Latency**: Small models respond faster. - **Cost**: Reduce serving compute costs. - **Edge**: Enable on-device inference. - **Efficiency**: Better than training small models from scratch. **Training Approaches** **Offline Distillation**: ``` 1. Train teacher model (or use pretrained) 2. Freeze teacher weights 3. Train student to match teacher Pro: Stable, simple Con: Fixed teacher, can't adapt ``` **Online Distillation**: ``` 1. Train teacher and student simultaneously 2. Student learns from evolving teacher 3. Sometimes mutual: both learn from each other Pro: Adaptive, can exceed static teacher Con: Complex, harder to optimize ``` **Self-Distillation**: ``` 1. Model distills to itself (deeper to shallower) 2. Or current model teaches previous version Pro: No separate teacher needed Con: Limited knowledge source ``` **Implementation** **Complete Training Loop**: ```python import torch import torch.nn as nn import torch.nn.functional as F class StudentTeacherTrainer: def __init__(self, teacher, student, temperature=4.0, alpha=0.5): self.teacher = teacher.eval() # Freeze teacher self.student = student self.temperature = temperature self.alpha = alpha self.optimizer = torch.optim.AdamW(student.parameters(), lr=1e-4) def distillation_loss(self, student_logits, teacher_logits, labels): # Soft loss (match teacher distribution) soft_targets = F.softmax(teacher_logits / self.temperature, dim=-1) soft_student = F.log_softmax(student_logits / self.temperature, dim=-1) soft_loss = F.kl_div(soft_student, soft_targets, reduction="batchmean") soft_loss *= self.temperature ** 2 # Hard loss (match true labels) hard_loss = F.cross_entropy(student_logits, labels) return self.alpha * hard_loss + (1 - self.alpha) * soft_loss def train_step(self, inputs, labels): # Teacher inference (no gradients) with torch.no_grad(): teacher_logits = self.teacher(inputs) # Student forward pass student_logits = self.student(inputs) # Compute loss loss = self.distillation_loss(student_logits, teacher_logits, labels) # Backprop self.optimizer.zero_grad() loss.backward() self.optimizer.step() return loss.item() ``` **Feature Distillation**: ```python class FeatureDistillationLoss(nn.Module): def __init__(self, student_dims, teacher_dims): super().__init__() # Projectors to match dimensions self.projectors = nn.ModuleList([ nn.Linear(s_dim, t_dim) for s_dim, t_dim in zip(student_dims, teacher_dims) ]) def forward(self, student_features, teacher_features): loss = 0 for proj, s_feat, t_feat in zip( self.projectors, student_features, teacher_features ): # Project student to teacher dimension s_proj = proj(s_feat) # MSE loss between features loss += F.mse_loss(s_proj, t_feat) return loss ``` **LLM Distillation** **Response-Based** (Common for LLMs): ```python def distill_llm(teacher, student, prompts): for prompt in prompts: # Teacher generates response with torch.no_grad(): teacher_response = teacher.generate( prompt, max_tokens=512, temperature=0.7 ) # Student learns to generate same response student_loss = student.forward( input_ids=prompt + teacher_response, labels=teacher_response # Predict teacher's tokens ) student_loss.backward() optimizer.step() ``` **Token-Level Matching**: ```python # Match next-token probabilities student_logits = student(input_ids).logits teacher_logits = teacher(input_ids).logits # KL divergence at each position loss = kl_div( log_softmax(student_logits / T), softmax(teacher_logits / T) ) * T² ``` **Model Size Guidelines** ``` Teacher Size | Student Size | Expected Retention ----------------|-----------------|-------------------- 70B parameters | 7B | 85-95% quality 7B parameters | 1.3B | 80-90% quality 1.3B parameters | 350M | 75-85% quality ``` **Architecture Choices**: ``` Option 1: Same architecture, fewer layers Option 2: Same architecture, smaller hidden dim Option 3: Different architecture entirely Best: Student architecture matches task needs ``` **Best Practices** ``` Practice | Recommendation ----------------------|---------------------------------- Data | Use teacher's training data if possible Temperature | Start with T=4, tune Training time | 1-3× normal epochs Learning rate | Lower than training from scratch Label smoothing | Often redundant with soft targets Intermediate layers | Match if architectures similar ``` Student-teacher learning is **the primary method for deploying powerful models efficiently** — by transferring knowledge from expensive-to-run teachers to compact students, organizations can deliver AI capabilities at a fraction of the inference cost.

Go deeper with CFSGPT

Get AI-powered deep-dives, save terms, and run advanced simulations — free account.

Create Free Account