Перейти к основному содержимому

Дистилляция знаний

Базовая идея

При использовании моделей глубокого обучения важна не только точность моделей, но и их вычислительная простота, обеспечивающая более быстрое построение прогнозов. Упрощение модели позволяет её применять на более простых вычислительных устройствах, таких как мобильный телефон, а также обрабатывать большее число пользовательских запросов за единицу времени.

Дистилляция знаний (knowledge distillation [1]) - процесс настройки простой модели, называемой учеником (student model) воспроизводить поведение уже настроенной сложной и точной модели, называемой учителем (teacher model). Впервые технология была развёрнуто описана в [2].

В базовом варианте дистилляция знаний для задачи классификации состоит из следующих шагов:

  1. Обучить простую модель f(x)f(x) восстанавливать вероятности классов сложной модели на трансферном датасете (transfer set).

Трансферный датасет может совпадать с тренировочным, на котором обучалась сложная модель, может составлять его подмножество, а может содержать и совсем другую выборку объектов.

Мотивация

Как правило, трансферный датасет мал, из-за чего возникают неоднозначности с обобщающей способностью - экстраполировать информацию об ограниченных наблюдениях можно по-разному. При этом известно, что сложная модель точна, т.е. обладает хорошей обобщающей способностью.

Поэтому, чтобы расширить объем получаемой информации с каждого объекта, простой модели даётся не просто информация о корректном классе (всего одно число), а предоставляется более полная информация о вероятностях каждого из CC классов по мнению сложной модели (CC чисел), что позволяет простой модели учиться быстрее, перенимая не только вероятную метку, но и "образ мышления" модели-учителя.

Рассмотрим классификацию изображений:

  • В стандартном обучении простая модель видит только изображение и верный класс, например, "собака". Переобучившись, она легко может научиться её путать с классом "медведь" и "лось" из-за цветовой схожести.

  • В дистилляции знаний простая модель видит, что "собака" обладает максимальным рейтингом, но также высоким рейтингом обладают классы "волк" и "шакал", а класс "медведь" и "лось" обладают, наоборот, низким рейтингом. Это позволяет простой модели перенимать "образ мышления" сложной модели, который обладает высокой обобщающей способностью.

Настройка простой модели

В стандартной процедуре обучения [2] простая модель выдаёт рейтинги f1(x),...fC(x)f_1(\mathbf{x}),...f_C(\mathbf{x}), которые трансформируются в вероятности классов через SoftMax-преобразование:

p1f,T(x)=ef1(x)/Tc=1Cefc(x/T),p2f,T(x)=ef2(x)/Tc=1Cefc(x)/T,    pCf,T(x)=efC(x)/Tc=1Cefc(x)/T,\begin{align*} & p^{f,T}_1(\mathbf{x}) = \frac{e^{f_1(\mathbf{x})/T}}{\sum_{c=1}^C e^{f_c(\mathbf{x}/T)}}, \\ & p^{f,T}_2(\mathbf{x}) = \frac{e^{f_2(\mathbf{x})/T}}{\sum_{c=1}^C e^{f_c(\mathbf{x})/T}}, \\ & \cdots \; \cdots \; \cdots\\ & p^{f,T}_C(\mathbf{x}) = \frac{e^{f_C(\mathbf{x})/T}}{\sum_{c=1}^C e^{f_c(\mathbf{x})/T}}, \\ \end{align*}

где T>0T>0 - гиперпараметр температуры, влияющий на степень сглаживания вероятностей. Чем температура выше, тем выходные вероятности получаются более сглаженные.

Настройка модели производится, используя стандартную кросс-энтропию с T=1T=1:

L(f(x),y)=c=1Cyclnpcf,1(x),yc=I{y=c}(1)\mathcal{L}(f(\mathbf{x}),y)=-\sum_{c=1}^C y_c\ln p^{f,1}_c(\mathbf{x}),\quad y_c=\mathbb{I}\{y=c\} \tag{1}

В дистилляции знаний есть уже настроенная сложная модель g()g(\cdot), выдающая вероятности классов [p1g,T(x),...,pCg,T(x)][p^{g,T}_1(\mathbf{x}), ..., p^{g,T}_C(\mathbf{x})]. Для простой модели эти вероятности представляют собой мягкую разметку (soft labels), к которой нужно приближать собственные прогнозы, для чего также используется кросс-энтропия:

Lkd(f(x)T,g)=c=1Cpcg,T(x)lnpcf,T(x)(2)\mathcal{L}_{kd}(f(\mathbf{x})|T,g)=-\sum_{c=1}^C p^{g,T}_c(\mathbf{x})\ln p^{f,T}_c(\mathbf{x}) \tag{2}

Используя (2), можно настраивать простую модель на даже неразмеченном трансферном датасете, поскольку разметку осуществляет модель-учитель. Если же трансферный датасет содержит метки истинных классов, то настроить простую модель можно точнее (в силу того, что модель-учитель всё же иногда ошибается), используя взвешенную сумму потерь (1) и (2):

L~(f(x),y)=L(f(x),y)+λT2Lkd(f(x)T,g),\tilde{\mathcal{L}}(f(\mathbf{x}),y) = \mathcal{L}(f(\mathbf{x}),y) + \lambda T^2 \mathcal{L}_{kd}(f(\mathbf{x})|T,g),

где λ>0\lambda>0 - гиперпараметр, отвечающий за силу дистилляции знаний.

В [2] рекомендуется стандартные потери L\mathcal{L} вычислять при T=1T=1, а дистиллированные потери Lkd\mathcal{L}_{kd} - с увеличенным T>1T>1. Это упрощает перенос "образа мышления" учительской модели моделью-учеником за счёт разглаживания вероятностей. В противном случае учительские вероятности часто будут оказываться слишком сконцентрированными вокруг предсказываемых классов, что затруднит дистилляцию знаний.

Поскольку градиент Lkd\mathcal{L}_{kd} по логитам потерь убывает по закону 1/T21/T^2 при увеличении TT (см. [1] и [2]), то, чтобы выровнять вклад потерь каждого типа, вторые потери домножаются на T2T^2.

Другие варианты применения

Выше мы рассмотрели перенос знаний сложной модели только с её последнего слоя, этот подход называется дистилляцией, основанной на откликах (responce-based distillation).

Однако часто простую модель обучают воспроизводить активации сложной модели на промежуточном слое. Для этого при настройке простой модели используется регуляризатор в виде квадрата L2L_2-нормы между картой признаков простой модели FK(x)F_K(x) и сложной модели GK(x)G_K(x) на слое KK:

L~(f(x),y)=L(f(x),y)+λFK(x)GK(x)2\tilde{\mathcal{L}}(f(\mathbf{x}),y) = \mathcal{L}(f(\mathbf{x}),y)+\lambda ||F_K(\mathbf{x})-G_K(\mathbf{x})||^2

Вместо квадрата L2L_2-нормы может использоваться и другая функция расхождения.

Этот подход называется дистилляция, основанная на признаках (feature-based distillation, [3]), поскольку активации сложной модели имеют смысл промежуточных признаков, которыми оперирует оставшаяся часть модели.

Также существует дистилляция, основанная на взаимодействии признаков (relation-based distillation, [3]), при которой простая модель учится воспроизводить такую же взаимосвязь между представлениями с разных слоёв ii и jj, что и сложная модель. Пусть Φ(u,v)\Phi(u,v) - функция взаимодействия между слоями. Например, это может быть матрица попарных корреляций между каждым признаком слоя ii с каждым признаком слоя jj. Тогда простая модель настраивается, используя следующую функцию потерь:

L~(f(x),y)=L(f(x),y)+λΦ(Fi(x),Fj(x))Φ(Gi(x),Gj(x))2\tilde{\mathcal{L}}(f(\mathbf{x}),y) = \mathcal{L}(f(\mathbf{x}),y)+\lambda ||\Phi(F_i(\mathbf{x}),F_j(\mathbf{x}))-\Phi(G_i(\mathbf{x}),G_j(\mathbf{x}))||^2

Здесь также вместо квадрата L2L_2-нормы может использоваться и другая функция расхождения.

Дистилляция от многих учителей

Простую нейросеть можно обучить качественнее, если использовать не одну, а сразу несколько сложных сетей [4]. Тогда можно будет собрать экспертизу с разных моделей, использующих разные подходы и принять более взвешенное решение.

Этого можно достичь разными способами:

  • Заставить простую модель приближать средний прогноз набора сложных моделей.

  • Заставить простую модель приближать прогноз каждой модели-учителя в отдельности (оптимизировать при этом сумму всех функций потерь)

  • Добавить внешний слой модели-студенту, в котором сделать столько выходов, сколько есть моделей-учителей. Заставить каждый выход приближать прогноз соответствующего учителя - это будет подстраивать прогноз модели-ученика под "вкусы" каждого учителя в отдельности. После настройки итоговый прогноз ученика строить как усреднение всех выходов.

Используя эти подходы, достаточно выразительная модель-студент может начать работать даже точнее, чем каждая из моделей-учителей.

Дистилляция, совмещённая с упрощением сложной модели

Существует большая разница в сложности между сложной и простой моделью, из-за чего простой модели может быть сложно перенять некоторые паттерны поведения сложной модели. В статье [5] удаётся повысить точность модель-ученика за счёт предварительного упрощения сложной модели за счёт её обрезки (network pruning).

Взаимное обучение

Методы, рассмотренные ранее, относятся к категории оффлайн-дистилляции, где процесс обучения строго последователен: сначала до полной сходимости обучается сложная модель-учитель g(x)g(\mathbf{x}), её веса фиксируются, и только затем начинается обучение модели-студента f(x)f(\mathbf{x}).

Если архитектура g(x)g(\mathbf{x}) значительно сложнее f(x)f(\mathbf{x}), то «образ мышления» учителя может оказаться слишком сложным для воспроизведения учеником, и модель-студент не сумеет в полной мере его воспроизвести.

В работе [6] был предложен подход взаимного обучения (Deep Mutual Learning, DML), который часто называют онлайн-дистилляцией. В этом сценарии модели f(x)f(\mathbf{x}) и g(x)g(\mathbf{x}) начинают обучение (со случайной инициализации) и учатся одновременно. При этом g(x)g(\mathbf{x}) не обязательно должна быть сложнее f(x)f(\mathbf{x}) — метод эффективно работает даже для моделей одинаковой архитектуры. Основная идея заключается в том, что две (или более) модели обучаются одновременно и совместно, обмениваясь знаниями на каждом шаге градиентного спуска.

Алгоритм для двух моделей

В процессе каждой итерации обучения обе модели получают один и тот же батч данных x\mathbf{x}. Для каждой модели вычисляется своя функция потерь, состоящая из стандартной кросс-энтропии с истинной разметкой yy и дистилляционных потерь, где роль «мягкой разметки» играют текущие предсказания модели-партнёра.

Для модели ff функция потерь имеет вид:

L~f(f(x),y)=L(f(x),y)+αLkd(f(x)T,g)\tilde{\mathcal{L}}_f(f(\mathbf{x}), y) = \mathcal{L}(f(\mathbf{x}), y) + \alpha \mathcal{L}_{kd}(f(\mathbf{x}) | T, g)

Для модели gg функция потерь симметрична:

L~g(g(x),y)=L(g(x),y)+αLkd(g(x)T,f)\tilde{\mathcal{L}}_g(g(\mathbf{x}), y) = \mathcal{L}(g(\mathbf{x}), y) + \alpha \mathcal{L}_{kd}(g(\mathbf{x}) | T, f)

В цикле до сходимости сэмплируются обучающие примеры (x,y)(\mathbf{x},y), по которым происходит обновление модели ff, потом модели gg.

В [6] предлагается брать α=1\alpha=1 и T=1T=1, а также расстояние Кульбака-Лейблера [7] между вероятностными распределениями двух моделей, которое эквивалентно указанной выше кросс-энтропии Lkd\mathcal{L}_{kd}.

Почему это работает?

Это работает как форма динамической регуляризации: каждая модель не просто минимизирует ошибку на данных, но и стремится найти такое решение, которое будет согласовано с мнением другой модели.

Хотя обе модели изначально «неопытны», они начинают обучение с разных случайных весов, а значит, находят разные статистические закономерности в данных. Обмениваясь вероятностями, они коллективно находят более устойчивые к переобучению решения.

Обобщение на KK моделей

Метод легко масштабируется на ансамбль из KK моделей {f1,f2,...,fK}\{f_1, f_2, ..., f_K\}. В этом случае для каждой модели fkf_k дистилляционные потери вычисляются как среднее расхождение с предсказаниями всех остальных участников ансамбля:

L~fk=L(fk(x),y)+αK1jkKLkd(fk(x)T,{fj}jk)\tilde{\mathcal{L}}_{f_k} = \mathcal{L}(f_k(\mathbf{x}), y) + \frac{\alpha}{K-1} \sum_{j \neq k}^K \mathcal{L}_{kd}(f_k(\mathbf{x}) | T, \{f_j\}_{j\ne k})

Такой подход позволяет еще сильнее повысить обобщающую способность каждой отдельной сети. После завершения обучения любую из моделей можно использовать независимо, при этом её вычислительная сложность останется прежней, а точность будет ближе к точности совместного ансамбля, состоящего из всех моделей.

Литература

  1. Wikipedia: Knowledge distillation.
  2. Hinton G., Vinyals O., Dean J. Distilling the knowledge in a neural network //arXiv preprint arXiv:1503.02531. – 2015.
  3. Gou J. et al. Knowledge distillation: A survey //International Journal of Computer Vision. – 2021. – Т. 129. – №. 6. – С. 1789-1819.
  4. Zuchniak K. Multi-teacher knowledge distillation as an effective method for compressing ensembles of neural networks //arXiv preprint arXiv:2302.07215. – 2023.
  5. Park J., No A. Prune your model before distill it //European Conference on Computer Vision. – Cham : Springer Nature Switzerland, 2022. – С. 120-136.
  6. Zhang Y. et al. Deep mutual learning //Proceedings of the IEEE conference on computer vision and pattern recognition. – 2018. – С. 4320-4328.
  7. Wikipeadia: Kullback–Leibler divergence.