# Function to compute distillation and hard-label loss def distillation_loss(student_logits, teacher_logits, true_labels, temperature, alpha): # Compute soft targets from teacher logits soft_targets = nn.functional.softmax(teacher_logits / temperature, dim=1) student_soft = nn.functional.log_softmax(student_logits / temperature, dim=1) # KL Divergence loss for distillation distill_loss = nn.functional.kl_div(student_soft, soft_targets, reduction='batchmean') * (temperature ** 2) # Cross-entropy loss for hard labels hard_loss = nn.CrossEntropyLoss()(student_logits, true_labels) # Combine losses loss = alpha * distill_loss + (1.0 - alpha) * hard_loss return loss