Skip to content
AITroveRead. Build. Understand.
Make this comfortable

Distillation soft targets, temperature and label balance

Last updated: 7 Oct 20265 min read
tutorial
AdvancedBy AITrove Editorial

A student can learn from a frozen teacher distribution and real labels, but the temperature, class map and loss reductions must be fixed before comparing candidates.

Freeze a trustworthy teacher

Distillation transfers a teacher’s behavior, including its blind spots. Evaluate the teacher on physical-receipt holdouts and rare defects before using it as supervision. Freeze its weights, put it in evaluation mode and record preprocessing, class order and revision. The student must see the same input semantics, even if its architecture is smaller. A teacher whose class index two means clipped receipt cannot supervise a student whose index two means blurred receipt. The quality task supplies a concrete class map and held-out identity rule.

Separate hard and soft objectives

The hard-label term compares unscaled student logits with human labels. The soft term compares the teacher and student distributions after both logits are divided by the same positive temperature. A common choice is teacher-to-student KL divergence with batch-mean reduction and a temperature-squared factor, which keeps gradient scale more comparable as temperature changes. Combine the terms with declared weights. Do not feed teacher probabilities into a loss expecting integer class IDs. Logit and loss contracts determine which tensor each term accepts.

Interpret the temperature

At temperature one, a confident teacher may assign almost no probability to alternatives. Raising temperature reveals relative preferences between non-top classes, but excessive smoothing can make the distribution close to uniform and uninformative. It does not improve teacher correctness. Compare temperatures using a development split and check class-specific student errors. The code below calculates one combined training step with a frozen teacher and verifies that only the student is updated. Teacher inference adds training cost but is absent from the final student inference path.

Guard against teacher error

A teacher can be confidently wrong on clipped edges or unusual capture devices. Inspect teacher disagreement with human labels by slice, and decide whether to downweight its soft loss for those records or obtain better labels. Never silently replace hard labels with teacher top classes. If teacher probabilities are cached, tie each row to receipt identity, preprocessing hash, teacher revision and class map; stale caches can train the student on the wrong image after an augmentation change. A student may surpass a teacher on some slices, but this must be measured rather than assumed.

Evaluate the intended deployment

Train a student from scratch with hard labels and another with the declared mixture, holding architecture, label set and training budget constant. Compare top-class accuracy, rare-defect recall, confidence quality, artifact size and actual target-device latency. Calibrate student probabilities separately after training if decisions use confidence; the distillation temperature is a training parameter, not automatically the serving calibration temperature. The applied release pairs the student comparison with edge constraints.

Implementation

python
import torch
from torch import nn
from torch.nn import functional as functional

torch.manual_seed(47)
receipt_features = torch.randn(7, 9)
quality_labels = torch.tensor([0, 2, 1, 2, 0, 1, 2])
teacher = nn.Sequential(nn.Linear(9, 28), nn.ReLU(), nn.Linear(28, 3))
student = nn.Sequential(nn.Linear(9, 11), nn.ReLU(), nn.Linear(11, 3))
teacher.eval()
for parameter in teacher.parameters():
    parameter.requires_grad_(False)
optimizer = torch.optim.AdamW(student.parameters(), lr=0.0008)
temperature = 2.3
with torch.no_grad():
    teacher_logits = teacher(receipt_features)
student_logits = student(receipt_features)
teacher_targets = functional.softmax(teacher_logits / temperature, dim=1)
student_log_probs = functional.log_softmax(student_logits / temperature, dim=1)
soft_loss = functional.kl_div(student_log_probs, teacher_targets,
                              reduction="batchmean") * temperature**2
hard_loss = functional.cross_entropy(student_logits, quality_labels)
training_loss = 0.38 * soft_loss + 0.62 * hard_loss
optimizer.zero_grad(set_to_none=True)
training_loss.backward()
optimizer.step()
assert all(parameter.grad is None for parameter in teacher.parameters())
assert torch.isfinite(training_loss)

Performance and operating cost

Training adds one teacher forward per student batch and stores student activations for backward propagation; a frozen teacher does not need an autograd graph. Cached teacher logits can reduce repeated teacher compute but require storage O(NC) for N records and C classes, plus strict identity and preprocessing versioning. Student inference cost depends only on the student graph. The temperature and KL operations are O(BC) per batch, usually small beside neural forward passes. Measure the final device path rather than inferring speed from parameter count.

Common Mistakes

  • Do not let gradients or training-mode statistics update the teacher.
  • Do not reuse the training temperature as an untested serving calibration value.
  • Do not accept teacher predictions as ground truth when they conflict with reviewed labels.

Read next

ai-data
deep-learning
Storage details