# put student model in train mode student_model.train() # train model for epoch in range(num_epochs): for batch in dataloader: # Prepare inputs input_ids = batch['input_ids'].to(device) attention_mask = batch['attention_mask'].to(device) labels = batch['labels'].to(device) # Disable gradient calculation for teacher model with torch.no_grad(): teacher_outputs = teacher_model(input_ids, attention_mask=attention_mask) teacher_logits = teacher_outputs.logits # Forward pass through the student model student_outputs = student_model(input_ids, attention_mask=attention_mask) student_logits = student_outputs.logits # Compute the distillation loss loss = distillation_loss(student_logits, teacher_logits, labels, temperature, alpha) # Backpropagation optimizer.zero_grad() loss.backward() optimizer.step() print(f"Epoch {epoch + 1} completed with loss: {loss.item()}") # Evaluate the teacher model teacher_accuracy, teacher_precision, teacher_recall, teacher_f1 = evaluate_model(teacher_model, test_dataloader, device) print(f"Teacher (test) - Accuracy: {teacher_accuracy:.4f}, Precision: {teacher_precision:.4f}, Recall: {teacher_recall:.4f}, F1 Score: {teacher_f1:.4f}") # Evaluate the student model student_accuracy, student_precision, student_recall, student_f1 = evaluate_model(student_model, test_dataloader, device) print(f"Student (test) - Accuracy: {student_accuracy:.4f}, Precision: {student_precision:.4f}, Recall: {student_recall:.4f}, F1 Score: {student_f1:.4f}") print("n") # put student model back into train mode student_model.train()