Ответ
Дистилляция знаний — это техника сжатия моделей, при которой небольшая, быстрая модель («студент») обучается имитировать поведение большой, сложной и точной модели («учитель») или их ансамбля. Цель — сохранить как можно больше производительности учителя при радикальном уменьшении вычислительных затрат и размера модели для деплоя на edge-устройствах или в условиях низкой задержки.
Ключевая идея: Студент обучается не только на «жёстких» метках (hard labels) исходного датасета, но и на «мягких» метках (soft labels) — распределениях вероятностей, которые выдает учитель. Эти мягкие метки содержат гораздо больше информации: они показывают, например, что изображение собаки больше похоже на волка, чем на самолет (отношения между классами), что помогает студенту обобщаться лучше.
Основные этапы:
- Обучение учителя: Большая модель обучается на исходных данных до высокой точности.
- Генерация мягких меток: Учитель делает предсказания на обучающих данных, получая распределения вероятностей по всем классам.
- Обучение студента: Малая модель обучается, минимизируя комбинированную функцию потерь:
- Loss между мягкими предсказаниями студента и учителя (Distillation Loss). Для этого используется «температура» (T) в функции softmax, чтобы сделать распределение учителя «мягче».
- Loss между предсказаниями студента и истинными метками (Student Loss).
Пример реализации на PyTorch:
import torch
import torch.nn as nn
import torch.nn.functional as F
class KnowledgeDistillationLoss(nn.Module):
def __init__(self, temperature=4.0, alpha=0.7):
super().__init__()
self.temperature = temperature
self.alpha = alpha # Вес для дистилляционного лосса
self.ce_loss = nn.CrossEntropyLoss()
self.kl_loss = nn.KLDivLoss(reduction='batchmean')
def forward(self, student_logits, teacher_logits, labels):
# 1. Лосс студента на истинных метках (hard loss)
hard_loss = self.ce_loss(student_logits, labels)
# 2. Лосс дистилляции (soft loss) с температурой
# Применяем softmax с температурой к логитам учителя и студента
soft_teacher = F.softmax(teacher_logits / self.temperature, dim=-1)
soft_student = F.log_softmax(student_logits / self.temperature, dim=-1)
# Используем KL-дивергенцию для сравнения распределений
soft_loss = self.kl_loss(soft_student, soft_teacher) * (self.temperature ** 2)
# 3. Комбинированный лосс
total_loss = (1.0 - self.alpha) * hard_loss + self.alpha * soft_loss
return total_loss
# Пример использования в цикле обучения
# teacher_model = ... (предобученная большая модель, в режиме eval())
# student_model = ... (малая модель для развертывания)
# criterion = KnowledgeDistillationLoss(temperature=4.0, alpha=0.7)
# for inputs, labels in dataloader:
# with torch.no_grad():
# teacher_logits = teacher_model(inputs)
# student_logits = student_model(inputs)
# loss = criterion(student_logits, teacher_logits, labels)
# loss.backward()
# optimizer.step()
Преимущества:
- Эффективное сжатие: Позволяет получить компактную модель, близкую по качеству к большой.
- Улучшенное обобщение: Студент часто превосходит модель, обученную только на жёстких метках, так как перенимает «знания» о сходстве классов.
- Практичность: Ансамбли тяжёлых моделей можно заменить одним лёгким студентом для продакшена.
Области применения: Сжатие BERT и других трансформеров (например, DistilBERT), развёртывание компьютерного зрения на мобильных устройствах, оптимизация моделей для IoT.