Зробити дистиляцію знань достатньо дешевою для масштабного застосування
Дистиляція знань — відома техніка в машинному навчанні, за якої меншу модель-учня навчають відтворювати продуктивність більшої моделі-вчителя. Завдяки нещодавній хвилі відкритих великих мовних моделей, таких як gpt-oss, Qwen, GLM або Kimi, вона знову стала провідною темою досліджень. Розгортання таких дуже великих моделей коштує дорого: нещодавня модель Kimi-K3 має 2,8 трильйона параметрів і потребує приблизно 3 ТБ VRAM лише для завантаження. Тому стискання їх до менших моделей і відновлення початкових можливостей за допомогою дистиляції знань стало стандартною практикою; такі компанії, як Nvidia (Nemotron 3 Puzzle 75B) та Multiverse Computing (Hypernova 60B), нещодавно випустили високоякісні стиснені моделі.
Саме етап дистиляції найбільше визначає кінцеву якість, але зазвичай він також є найдорожчою частиною конвеєра. Одночасне завантаження моделі-вчителя й моделі-учня та створення розподілу ймовірностей за всім словником для кожного токена потребують величезного обсягу VRAM і зазвичай можливі лише за наявності сотень GPU та ретельно продуманих стратегій тензорного паралелізму. У нашій останній статті Efficient Knowledge Distillation for LLMs: Offline Top-K Logits and a Fused Chunked KL Loss ми розв’язуємо цю проблему за допомогою двох системних змін: одноразового кешування top-K логітів учителя, щоб учитель більше ніколи не мусив перебувати в пам’яті разом з учнем, і нового, економного щодо пам’яті лосу KL-дивергенції, який уникає матеріалізації повної матриці розміром словник × довжина послідовності. Це дає змогу скоротити використання VRAM значно нижче рівня, якого досягають стандартні реалізації в бібліотеках на кшталт PyTorch або NVIDIA Megatron-Bridge. Разом ці дві зміни знижують вартість навчання настільки, що відновлення моделей із довгим контекстом стає можливим на одному GPU, а масштабні експерименти — практичними.
Чому відновлення за допомогою дистиляції є дорогим
Стандартна конфігурація — онлайн-дистиляція із використанням лосу дивергенції Кульбака—Лейблера (лосу KL) — одночасно підтримує завантаженими і вчителя, і учня. На кожному кроці навчання вчитель виконує повний прямий прохід, щоб створити свій вихідний розподіл, а учня навчають відтворювати його. Це найвиразніша конфігурація, оскільки доступний повний розподіл учителя, але водночас вона потребує найбільше пам’яті та обчислень: для кожної позиції токена потрібно утримувати два тензори повного словника, а вчителя доводиться перераховувати на кожному кроці, хоча його поведінка протягом навчання не змінюється.
Розглянемо практичний приклад: gpt-oss-120b має словник із 201 088 токенів. За довжини послідовності 32K і розміру пакета 4 сам тензор імовірностей учителя має форму 4 × 201,088 × 32,768; у форматі bfloat16 це вже близько 50 ГБ VRAM для одного тензора. Якщо додати градієнти, активації, ваги моделі та стани оптимізатора, одна ітерація дистиляції може досягати приблизно 250 ГБ VRAM — більше, ніж може надати навіть GPU H200 або B200. У цій публікації ми показуємо, що переформулювання лосу KL для обробки даних частинами знижує ці витрати майже до нуля.
Щільний KL різко зростає приблизно до 250 ГБ — вище за ємність одного H200 (141 ГБ). Комбінований частковий лос не створює такого піку, а його максимум становить близько 128 ГБ. Джерело: рисунок 1 статті.
Дві системні зміни
Офлайн-дистиляція. Замість перерахунку вчителя на кожному кроці ми обчислюємо його вихід один раз, кешуємо 100 найбільш імовірних токенів для кожної позиції та навчаємо учня на основі цього кешу. Учителю більше не потрібно перебувати в пам’яті під час навчання, і після створення кешу його не потрібно запускати повторно, тож той самий кеш можна використовувати для багатьох абляцій.
Комбінований частковий лос KL. Щоб зрозуміти, чому сам лос є дорогим, уявімо, що саме він створює: для кожної позиції токена в послідовності та кожного слова у словнику лосу потрібне число, яке описує, наскільки передбачення учня не збігається з передбаченням учителя. Якщо подати це у вигляді сітки, маємо один рядок на кожен елемент словника й один стовпчик на кожну позицію послідовності; для словника зі 100 тисячами й більше слів і довгої послідовності ця сітка величезна, а стандартний спосіб обчислення лосу KL створює її повністю, перш ніж зможе отримати хоча б одне число.
Ми порівнюємо три способи обчислення цього самого лосу, математично еквівалентні між собою:
- Щільний KL — підхід із підручника. Він повторно створює повну щільну сітку ймовірностей учителя з кешованих top-100 логітів і порівнює її з власною щільною сіткою логарифмів імовірностей учня. Ця версія найближча до того, як уже працює онлайн-дистиляція, тому ми використовуємо її як базову для перевірки коректності, але вона двічі зберігає в пам’яті повну сітку словник × послідовність.
- KL із частковою обробкою прямого проходу зберігає розрідженість учителя (лише кешовані top-100 логітів для кожної позиції, без розгортання в щільну сітку) і обчислює лос частинами, по одному фрагменту позицій послідовності за раз. Це усуває щільного вчителя та щільне порівняння і, як виявилося, є найшвидшим із трьох методів у наших тестах. Однак у нього залишається одна сліпа зона: власні логіти учня, тобто сітка, створена вихідним шаром моделі, усе ще обчислюються повністю та зберігаються для зворотного проходу, тому обсяг пам’яті все ще стрімко зростає разом із довжиною послідовності.
- Комбінований частковий KL — наш основний внесок — робить ще один крок і безпосередньо об’єднує вихідну проєкцію моделі з обчисленням лосу. Він взагалі не створює повної сітки логітів учня: обробляє по одному фрагменту послідовності за раз від початку до кінця, проєктує приховані стани в логіти для цього фрагмента, додає результат до поточного лосу та відкидає фрагмент перед переходом до наступного. Під час зворотного проходу кожен фрагмент перераховується на льоту, а не зберігається. Ціна цього — подвійне виконання проєкції: один раз у прямому проході й один раз у зворотному. Натомість пікове споживання пам’яті зростає лише лінійно разом із довжиною послідовності, а не стрибкоподібно через повний розмір словника × послідовності.
Наведена нижче GIF-анімація показує різницю між щільним і комбінованим частковим підходами: один створює всю сітку порівняння та утримує її, інший створює й відкидає по одному фрагменту за раз, тому пам’ять ніколи не зростає понад розмір одного фрагмента.
Ми опублікували реалізацію часткового лосу у відкритому доступі: github.com/CompactifAI/Full-Chunked-KL-Loss
Що це змінює на практиці
У таблиці нижче напряму порівнюються всі чотири конфігурації: онлайн-дистиляція та три щойно описані реалізації офлайн-лосу. Порівнюючи їх на одному GPU H200 із Llama 3.1 8B Instruct як учителем і моделлю Llama на 3,2 млрд параметрів як учнем, за контексту в 8K токенів, усі чотири досягають майже ідентичного лосу навчання, хоча офлайн-запуски навчаються лише на кешованих top-100 логітах для кожного токена.
| Метод (контекст 8K, один H200) | Пікове споживання пам’яті | Час ітерації | Пропускна здатність |
|---|---|---|---|
| Онлайн-дистиляція | 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 |
Криві лосу майже повністю збігаються для всіх чотирьох методів, підтверджуючи, що офлайн-дистиляція з кешованими top-100 логітами не втрачає якості порівняно з онлайн-дистиляцією. Джерело: рисунок 2 статті. За такої довжини послідовності комбінований частковий лос ще не є найшвидшим варіантом: додаткова проєкція під час зворотного проходу дещо знижує швидкість, але його справжня перевага проявляється зі збільшенням довжини контексту, що демонструє наступний розділ.
Масштабування до довгих контекстів
Щоб виразніше побачити закономірність масштабування, ми провели окремий тест на мережі вихідної проєкції (без тіла трансформера, лише ядро лосу). За 32K токенів пікове споживання пам’яті зменшується з 85.2 ГіБ для щільного лосу до 5.45 ГіБ для повністю часткового варіанта — у 15,6 раза, а щільний лос повністю виходить з ладу, починаючи з 64K токенів. За 256K токенів повністю частковий лос використовує 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,2 млрд параметрів, зберігає більшу частину точності вчителя на BoolQ і HellaSwag, а на MMLU відстає від нього приблизно на дев’ять пунктів, маючи менш ніж половину кількості параметрів.
Модель-учень зберігає більшу частину точності вчителя на короткому контексті, маючи менш ніж половину його розміру. Джерело: рисунок 6 статті.
Ця робота є частиною постійних досліджень Multiverse Computing, спрямованих на те, щоб зробити дистиляцію та відновлення практичними для масштабного запуску — не просто одноразовим рецептом, а процесом, над яким команди можуть дешево експериментувати. У статті також розглянуто додаткові абляції, зокрема вплив вибору функції лосу та пакування послідовностей на якість відновлення.
Потрібні повні технічні деталі, зокрема градієнт у замкненій формі для комбінованого часткового лосу та повна конфігурація навчання? Прочитайте повну статтю або зв’яжіться з нашою командою, щоб обговорити застосування цього підходу до власних конвеєрів дистиляції.
Ми також опублікували реалізацію часткового лосу у відкритому доступі: github.com/CompactifAI/Full-Chunked-KL-Loss
Моделі, згадані в цій статті 6
Статті, згадані в цій статті 1
Колекції, згадані в цій статті 1
Спільнота
· Зареєструйтеся або увійдіть, щоб залишити коментар
Моделі, згадані в цій статті 6
Статті, згадані в цій статті 1
Колекції, згадані в цій статті 1
Перекладено автоматично з англійської. Оригінал статті — за посиланням нижче.

