Мой блог
Как сделать дистилляцию LLM дешевле: эффективные методы обучения

Эффективная дистилляция знаний: как масштабировать обучение моделей
Дистилляция знаний — процесс обучения менее мощной модели-ученика для имитации поведения более крупной «учительской» модели — является широко известной и фундаментальной техникой в машинном обучении. С недавним всплеском популярности открытых больших языковых моделей (LLM), таких как gpt-oss, Qwen, GLM или Kimi, эта методика вновь стала центральной темой для исследований в индустрии. Актуальность обусловлена практической сложностью развертывания гигантских систем: например, недавняя модель Kimi-K3 оперирует 2,8 триллионами параметров и требует примерно 3 ТБ видеопамяти (VRAM) только для своей загрузки. Сжатие таких моделей с последующим восстановлением их исходных способностей через дистилляцию знаний стало общепринятым стандартом. Компании, такие как Nvidia (с моделью Nemotron 3 Puzzle 75B) и Multiverse Computing (с Hypernova 60B), недавно выпустили высококачественные сжатые модели, демонстрирующие возможности этого подхода.
Тем не менее, этап дистилляции зачастую является наиболее дорогостоящей частью pipeline’а. Поддержание в памяти одновременно учителя и ученика, а также генерация полного распределения вероятностей по всему словарю для каждого токена требуют колоссальных объемов VRAM. Это обычно под силу лишь кластерам из сотен GPU при условии использования тщательных стратегий тензорного параллелизма. В нашей недавней статье «Efficient Knowledge Distillation for LLMs: Offline Top-K Logits and a Fused Chunked KL Loss» мы решаем эту проблему с помощью двух системных инноваций: кеширования топ-K логитов учителя, что исключает необходимость удержания учителя в оперативной памяти во время обучения, и нового метода вычисления KL-дивергенции, эффективного с точки зрения потребления памяти. Этот метод позволяет избежать формирования полной матрицы (размер словаря × длина последовательности), снижая потребление VRAM до значений, значительно меньших, чем у стандартных реализаций в таких библиотеках, как PyTorch или NVIDIA Megatron-Bridge. В совокупности эти два изменения сокращают затраты на обучение настолько, что становится возможным «заживление» (healing) длинного контекста на одном GPU, а крупномасштабные эксперименты становятся практически доступными.
Ограничения стандартной онлайн-дистилляции
Традиционный подход, называемый онлайн-дистилляцией с использованием функции потерь Кульбака-Лейблера (KL loss), требует одновременной загрузки учителя и ученика в память. На каждом шаге обучения учитель выполняет полный проход, чтобы сгенерировать свое выходное распределение, которому должен следовать ученик. Хотя это наиболее выразительный метод, поскольку доступно полное распределение учителя, он крайне ресурсоемок: для каждой позиции токена необходимо хранить два тензора размером с весь словарь, а учитель должен пересчитываться на каждом шаге, несмотря на то, что его поведение не меняется в процессе обучения. В качестве практического примера рассмотрим модель gpt-oss-120b со словарем в 201 088 токенов. При длине последовательности 32K и размере батча 4, один лишь тензор вероятностей учителя имеет размерность 4 × 201 088 × 32 768. В формате bfloat16 это уже около 50 ГБ VRAM только для одного тензора. При добавлении градиентов, активаций, весов модели и состояний оптимизатора одна итерация дистилляции может достигать пикового потребления в 250 ГБ VRAM, что превышает возможности даже таких ускорителей, как H200 или B200.
Офлайн-дистилляция и работа с чанками
В нашей работе мы доказываем, что переформулировка KL-loss для обработки данных по частям (чанками) позволяет почти нивелировать эти затраты. Плотный (dense) KL-loss показывает всплеск потребления памяти до 250 ГБ, что выше емкости в 141 ГБ у одного H200. Напротив, нашFused-chunked метод позволяет никогда не создавать подобный всплеск, ограничивая пиковое потребление примерно 128 ГБ. При офлайн-дистилляции вместо пересчета учителя на каждом шаге мы вычисляем его выходы один раз, кешируем топ-100 наиболее вероятных токенов для каждой позиции и обучаем ученика против этого кеша. Учитель никогда не должен находиться в памяти во время обучения и не требует повторного запуска, а тот же кеш может быть повторно использован для множества различных абляционных исследований.
Для понимания того, почему сама функция потерь столь затратна, представьте, что она фактически строит: для каждой позиции токена в последовательности и каждого слова в словаре необходимо число, описывающее, насколько предсказание ученика расходится с предсказанием учителя. Для словаря в 100 000+ слов и длинной последовательности эта сетка становится гигантской, а стандартный способ вычисления KL-loss строит всю эту структуру целиком, прежде чем получить хотя бы одно число. Мы сравнили три математически эквивалентных способа вычисления этой функции:
- Dense KL (эталонный подход): воссоздает полную, плотную сетку вероятностей учителя из кешированных топ-100 логитов и сравнивает её с собственной плотной сеткой логитов ученика. Это версия, наиболее близкая к тому, как работает онлайн-дистилляция, поэтому мы используем её как базовую линию для проверки корректности. Однако она удерживает полную сетку «словарь × последовательность» в памяти в двойном размере.
- Forward-chunked KL: сохраняет разреженность учителя (только кешированные топ-100 логитов для каждой позиции, никогда не расширяясь до плотной сетки) и вычисляет потери по частям, один слайс последовательности за раз. Это удаляет необходимость в плотном учителе и плотном сравнении, что делает метод самым быстрым среди трех в наших бенчмарках. Однако у него есть «слепое пятно»: логиты ученика (сетка, создаваемая выходным слоем модели) по-прежнему вычисляются целиком и удерживаются для обратного прохода, поэтому потребление памяти все еще круто растет с увеличением длины последовательности.
- Fused chunked KL (наш основной вклад): идет на шаг дальше и встраивает проекцию выходных данных модели непосредственно в процесс вычисления потерь. Этот метод вообще не создает полную сетку логитов ученика. Он обрабатывает один чанк последовательности за раз «от и до», проектируя скрытые состояния в логиты для этого чанка, складывая результат в текущее значение потерь и отбрасывая чанк перед переходом к следующему. Обратный проход пересчитывает каждый чанк на лету вместо его хранения. Цена этого — выполнение проекции дважды (один раз вперед, один раз назад), но взамен пиковая память растет лишь линейно с длиной последовательности, вместо резких всплесков при учете всего размера словаря.
Результаты и масштабирование
Мы провели прямое сравнение всех четырех методов на одном H200 GPU при использовании Llama 3.1 8B Instruct в качестве учителя и 3.2B модели Llama в качестве ученика с контекстом 8K токенов. Все четыре метода достигают практически идентичных потерь при обучении, подтверждая, что офлайн-дистилляция с кешированными топ-100 логитами не приводит к потере точности относительно онлайн-дистилляции. На данной длине последовательности fused chunked loss еще не является самым быстрым, так как его дополнительная проекция в обратном проходе требует немного больше времени, но его реальное преимущество проявляется по мере роста длины контекста. Изолированный бенчмарк на тестовой сети (без тела трансформера, только ядро потерь) показал, что на 32K токенов пиковая память падает с 85,2 ГиБ при dense loss до 5,45 ГиБ при fully chunked версии — это 15,6-кратное сокращение. Dense loss терпит неудачу при 64K токенов и выше. На 256K токенов fully chunked loss использует 11,6 ГиБ против 134,2 ГиБ у следующего по эффективности чанкового варианта, при этом он примерно в 3,3 раза быстрее за итерацию. Дистилляция модели GPT-OSS 20B при контексте 32 768 токенов позволила сократить инфраструктуру с четырех GPU-узлов до одного. Время шага упало с 57,0 до 12,23 секунд (примерно в 5 раз быстрее), а пропускная способность на один GPU выросла с 74,2 до 345,7 TFLOP/s. Эффективная офлайн-установка — это то, что сделало крупномасштабную дистилляцию экономически оправданной. В итоге компактный ученик, дистиллированный из Llama 3.1 8B Instruct до 3.2B параметров, сохраняет большинство показателей точности учителя на тестах BoolQ и HellaSwag, оставаясь в пределах девяти баллов от него на MMLU при менее чем половине количества параметров. Данная работа является частью текущих исследований Multiverse Computing, направленных на то, чтобы сделать дистилляцию и «заживление» моделей практически применимыми в промышленном масштабе, а не просто разовыми рецептами, позволяя командам проводить итерации дешево. Мы также опубликовали реализацию chunked-loss с открытым исходным кодом: github.com/CompactifAI/Full-Chunked-KL-Loss. Статья также охватывает дополнительные абляции, такие как влияние выбора функции потерь и упаковки последовательностей на качество восстановления. Полные технические детали, включая замкнутый градиент, стоящий за fused chunked loss, можно найти в оригинальной статье.
Источник: huggingface.co

