Загрузка 0
ПОДЕЛИТЬСЯ

Мой блог

Листай вниз

Убыстрение и удешевление дистилляции знаний для больших языковых моделей

Убыстрение и удешевление дистилляции знаний для больших языковых моделей

Дистилляция знаний и масштабные языковые модели

Дистилляция знаний, заключающаяся в обучении меньшей «модели-ученика» для достижения эффективности крупной «модели-учителя», представляет собой хорошо изученный метод в области машинного обучения. С недавним появлением волны открытых больших языковых моделей, таких как gpt-oss, Qwen, GLM или Kimi, эта тема вновь оказалась в центре внимания исследователей. Развертывание столь масштабных архитектур сопряжено с серьезными затратами: например, недавняя модель Kimi-K3 насчитывает 2,8 триллиона параметров и требует около 3 терабайт видеопамяти только для первичной загрузки. Сжатие таких систем в более компактные варианты с последующим восстановлением исходных возможностей через дистиллят стало стандартной практикой. Компании вроде Nvidia с их Nemotron 3 Puzzle 75B или Multiverse Computing с Hypernova 60B недавно представили высококачественные сжатые версии нейросетей.

Процесс дистилляции во многом определяет итоговое качество продукта, однако именно он традиционно оказывается наиболее затратным этапом конвейера. Одновременное удержание в памяти учителя и ученика, а также генерация вероятностного распределения по всему словарю для каждого отдельного токена требуют колоссальных объемов VRAM. На практике это обычно осуществимо лишь при задействовании сотен графических процессоров и применении сложных стратегий тензорного параллелизма. Наша недавняя научная работа «Эффективная дистилляция знаний для LLM: офлайн-логиты Top-K и объединенная блочная потеря KL» решает эту проблему с помощью двух системных инноваций: однократного кэширования топ-K логитов учителя, благодаря чему модель-учитель вообще не обязана находиться в памяти рядом с учеником, и новой функции потерь на основе дивергенции Кульбака — Лейблера, экономящей память и полностью исключающей необходимость материализации громоздкой матрицы размера «размер словаря на длину последовательности».

Преодоление аппаратных ограничений стандартной дистилляции

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

В качестве наглядного примера возьмем модель gpt-oss-120b со словарем в 201 088 токенов. При длине последовательности 32K и размере батча 4 один только тензор вероятностей учителя имеет форму 4 × 201 088 × 32 768. В формате bfloat16 это составляет примерно 50 гигабайт VRAM исключительно для одного тензора. Если прибавить градиенты, активации, веса модели и состояния оптимизатора, пиковое потребление памяти на одну итерацию дистилляции может достигать 250 гигабайт — показатель, превышающий емкость даже современных ускорителей H200 или B200.

В нашем исследовании демонстрируется, что переработка функции потерь KL для обработки данных блоками сокращает эти издержки почти до нуля. Плотная функция KL дает скачок потребления примерно до 250 ГБ, что выше лимита в 141 ГБ для одиночного H200. Объединенная блочная функция потерь предотвращает этот всплеск, удерживая пиковые показатели на отметке около 128 ГБ, как показано на рисунке 1 в оригинальной работе.

Офлайн-подход и оптимизация вычислений

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

Чтобы понять затратность традиционной функции потерь, представьте структуру, которую она создает: для каждой позиции токена в последовательности и каждого слова из словаря требуется числовое описание степени расхождения между предсказанием ученика и учителя. В виде сетки это дает по одной строке на каждый элемент словаря и по одному столбцу на каждую позицию последовательности. Для словаря объемом более 100 тысяч слов и длинного контекста такая сетка оказывается колоссальной, а стандартный способ расчета KL формирует ее целиком до получения хотя бы одного числового результата.

Мы сопоставили три математически эквивалентных метода вычисления этой потери:

  • Плотный KL (Dense KL) — классический подход из учебников. Он восстанавливает полную плотную сетку вероятностей учителя из кэшированных топ-100 логитов и сопоставляет ее с собственной плотной сеткой лог-вероятностей ученика. Этот вариант ближе всего к классической онлайн-дистилляции, поэтому мы используем его в качестве базовой линии корректности, но он удерживает в памяти двойную громоздкую сетку «словарь на последовательность».
  • Блочный KL вперед (Forward-chunked KL) — сохраняет разреженность учителя (используя только кэшированные топ-100 логитов без расширения до плотной сетки) и считает потери по частям, срезами позиций последовательности. Это исключает плотного учителя и плотное сравнение, оказываясь самым быстрым методом в наших тестах. Тем не менее, сохраняется узкое место: собственные логиты ученика вычисляются полностью и удерживаются для обратного прохода, из-за чего память по-прежнему резко растет с увеличением длины контекста.
  • Объединенный блочный KL (Fused chunked KL) — наше ключевое достижение. Метод интегрирует проекцию выхода модели непосредственно в вычисление потерь. Он полностью избегает формирования сетки полных логитов ученика: обработка идет по одному блоку последовательности за раз от начала до конца, проецируя скрытые состояния в логиты для конкретного блока, сворачивая результат в текущее значение потерь и отбрасывая блок перед переходом к следующему. Обратный проход пересчитывает каждый блок на лету вместо его сохранения. Цена вопроса — двукратное выполнение проекции (вперед и назад), но взамен пиковое потребление памяти растет лишь линейно относительно длины последовательности, а не взмывает вверх пропорционально произведению словаря на длину.

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

Результаты бенчмарков и масштабирование

Мы открыли исходный код реализации блочных потерь в публичном репозитории. Таблица ниже сопоставляет все четыре конфигурации: онлайн-дистилляцию и три описанных офлайн-метода потерь. Тестирование проводилось на единичном ускорителе H200 с моделью Llama 3.1 8B Instruct в роли учителя и версией Llama с 3,2 млрд параметров в качестве ученика при контексте в 8 тысяч токенов. Все четыре варианта демонстрируют практически идентичные кривые потерь, подтверждая, что офлайн-дистилляция с кэшированием топ-100 логитов не уступает по качеству онлайн-подходу.

  • Онлайн-дистилляция: 102.8 ГБ памяти, 25.9 с на итерацию, 237 TFLOP/s.
  • Офлайн, плотный KL: 78.3 ГБ памяти, 18.5 с на итерацию, 331 TFLOP/s.
  • Офлайн, блочный KL вперед: 61.8 ГБ памяти, 18.4 с на итерацию, 335 TFLOP/s.
  • Офлайн, объединенный блочный KL: 58.3 ГБ памяти, 20.2 с на итерацию, 304 TFLOP/s.

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

Для более четкой демонстрации масштабирования мы провели изолированный бенчмарк на тестовой сети выходной проекции без тела трансформера (исключительно ядро потерь). На 32 тысячах токенов пиковое потребление памяти падает с 85,2 ГиБ у плотных потерь до 5,45 ГиБ у полностью блочной версии — 15,6-кратное сокращение, тогда как плотный метод полностью аварийно завершает работу на отметке от 64 тысяч токенов и выше. При длине в 256 тысяч токенов полностью блочные потери задействуют 11,6 ГиБ против 134,2 ГиБ у ближайшего конкурента среди блочных вариантов, обеспечивая при этом примерно в 3,3 раза большую скорость на одну итерацию.

Итоги и практическое применение

Дистилляция модели GPT-OSS 20B с контекстом в 32 768 токенов с использованием нашей памяти позволила сократить инфраструктуру с четырех узлов графических процессоров до одного. Время шага сократилось с 57,0 до 12,23 секунды (примерно пятикратное ускорение), а пропускная способность на один GPU возросла с 74,2 до 345,7 TFLOP/s. Эффективная офлайн-конфигурация сделала крупномасштабную кампанию по дистилляции экономически жизнеспособной.

Полученный компактный ученик, дистиллированный из Llama 3.1 8B Instruct до примерно 3,2 млрд параметров, сохраняет большую часть точности учителя на бенчмарках BoolQ и HellaSwag, отставая лишь в пределах девяти пунктов на MMLU при менее чем половинном количестве параметров (рисунок 6 в статье).

Данная работа является частью продолжающихся исследований компании Multiverse Computing, направленных на превращение процессов дистилляции и дообучения в практически применимые инструменты масштабирования, доступные для регулярных итераций командами разработчиков. В публикации также рассматриваются дополнительные аспекты абляции, включая влияние выбора функции потерь и упаковки последовательностей на качество восстановления. Полные технические детали, включая аналитический градиент для объединенных блочных потерь и конфигурацию обучения, доступны в оригинальной статье. Вы также можете связаться с нашей командой для обсуждения внедрения этого метода в ваши собственные конвейеры дистилляции. Открытый исходный код реализации блочных потерь опубликован по адресу github.com/CompactifAI/Full-Chunked-KL-Loss.

Источник: huggingface.co

Оставить комментарий

Ваш адрес email не будет опубликован. Обязательные поля помечены *

01.
На платформе MonsterInsights