Дестилацията на знанията да стане достатъчно евтина за мащабно прилагане
Дестилацията на знания, при която се обучава по-малък ученически модел да постигне производителността на по-голям учителски модел, е добре позната техника в машинното обучение. С последната вълна от отворени големи езикови модели, като gpt-oss, Qwen, GLM или Kimi, тя отново се превърна в основна изследователска тема. Внедряването на тези много големи модели е скъпо: скорошният модел Kimi-K3 има 2,8 трилиона параметъра и се нуждае от приблизително 3 TB VRAM само за да бъде зареден. Затова компресирането им в по-малки модели и възстановяването на първоначалните им способности чрез дестилация на знания се превърнаха в стандартна практика, като компании като Nvidia (Nemotron 3 Puzzle 75B) или Multiverse Computing (Hypernova 60B) наскоро пуснаха висококачествени компресирани модели.
Именно стъпката на дестилация определя по-голямата част от крайното качество, но обикновено тя е и най-скъпата част от процеса. Поддържането на учителския и ученическия модел заредени едновременно и създаването на вероятностно разпределение върху целия речник за всеки токен изискват огромно количество VRAM и обикновено са осъществими само със стотици графични процесори и внимателно планирани стратегии за тензорен паралелизъм. Последната ни статия, Ефективна дестилация на знания за LLM: логити Offline Top-K и слята Chunked KL загуба, решава този проблем чрез две системни промени: кеширане на top-K логитите на учителя еднократно, така че учителят никога да не трябва да се намира в паметта заедно с ученика, и нова, икономична по отношение на паметта загуба на KL-дивергенция, която избягва материализирането на матрицата с размер речник × дължина на последователността, намалявайки използването на VRAM далеч под нивата, постигани от стандартните реализации в библиотеки като PyTorch или NVIDIA Megatron-Bridge. Заедно тези две промени намаляват разходите за обучение достатъчно, за да направят възможно възстановяването на дълъг контекст на един графичен процесор и да направят експериментите в голям мащаб практически осъществими.
Защо възстановяването чрез дестилация е скъпо
Стандартната конфигурация — онлайн дестилация с използване на загубата на дивергенция на Kullback-Leibler (KL загуба) — поддържа учителския и ученическия модел заредени едновременно. При всяка стъпка на обучение учителят извършва пълен проход напред, за да създаде изходното си разпределение, а ученикът се обучава да го наподобява. Това е най-изразителната конфигурация, тъй като е налично пълното разпределение на учителя, но също така е най-интензивната по отношение на паметта и изчисленията: за всяка позиция на токен трябва да се съхраняват два тензора върху целия речник, а учителят трябва да се преизчислява при всяка отделна стъпка, въпреки че поведението му не се променя в рамките на обучението.
Като практически пример, gpt-oss-120b има речник от 201 088 токена. При дължина на последователността 32K и размер на партидата 4 само тензорът на вероятностите на учителя има размер 4 × 201,088 × 32,768; при bfloat16 това вече е около 50 GB VRAM за един тензор. Като добавим градиенти, активации, тегла на модела и състояния на оптимизатора, една итерация на обучение с дестилация може да достигне пиково приблизително 250 GB VRAM — повече, отколкото може да предостави дори графичен процесор H200 или B200. В тази публикация показваме, че преформулирането на KL загубата за обработка на данните на части намалява тези разходи почти до нула.
Плътната KL загуба достига приблизително 250 GB, над капацитета от 141 GB на един H200. Слятата Chunked загуба никога не създава такъв пик и достига максимум около 128 GB. Източник: Фигура 1 от статията.
Две системни промени
Offline дестилация. Вместо да преизчисляваме учителя при всяка стъпка, изчисляваме изхода му веднъж, кешираме 100-те най-вероятни токена за всяка позиция и обучаваме ученика спрямо този кеш. Учителят никога не трябва да се намира в паметта по време на обучението и не е необходимо да се стартира отново, след като кешът е наличен, така че един и същ кеш може да се използва повторно при множество аблации.
Слята, Chunked KL загуба. За да разберем защо самата загуба е скъпа, нека си представим какво всъщност създава тя: за всяка позиция на токен в дадена последователност и за всяка дума в речника загубата се нуждае от число, описващо доколко предсказанието на ученика се разминава с това на учителя. Представено като мрежа, това е по един ред за всеки запис в речника и по една колона за всяка позиция в последователността; при речник от над 100K думи и дълга последователност тази мрежа е огромна, а стандартният начин за изчисляване на KL загуба изгражда цялата мрежа, преди да може да произведе дори едно число.
Сравняваме три начина за изчисляване на същата загуба, които са математически еквивалентни:
- Плътна KL загуба е подходът от учебниците. Тя възстановява пълна, плътна мрежа от вероятности на учителя от кешираните top-100 логити и я сравнява със собствената плътна мрежа от логаритмични вероятности на ученика. Това е версията, която е най-близка до начина, по който вече работи онлайн дестилацията, затова я използваме като базова линия за коректност, но тя съхранява в паметта цялата мрежа речник × последователност, при това два пъти.
- Forward-chunked KL запазва разредения учител (само кешираните му top-100 логити за всяка позиция, без разширяване до плътна мрежа) и изчислява загубата по части, по един отрязък от позиции в последователността. Това премахва плътния учител и плътното сравнение и се оказва най-бързият от трите метода в нашите тестове. Все пак има една сляпа зона: собствените логити на ученика — мрежата, създадена от изходния слой на модела — все още се изчисляват изцяло и се запазват за обратния проход, така че паметта продължава да нараства рязко с дължината на последователността.
- Слята Chunked KL, основният ни принос, прави още една крачка напред и слива директно изходната проекция на модела с изчисляването на загубата. Тя изобщо не създава пълната мрежа от логити на ученика: обработва по един отрязък от последователността от начало до край, проектира скритите състояния към логити за този отрязък, включва резултата в текущата загуба и изхвърля отрязъка, преди да премине към следващия. При обратния проход всеки отрязък се преизчислява в движение, вместо да се съхранява. Цената е тази проекция да се извършва два пъти — веднъж напред и веднъж при обратния проход — но в замяна пиковото използване на паметта нараства само линейно с дължината на последователността, вместо да достига пик заради пълния размер речник × последователност.
GIF-ът по-долу показва разликата между плътния и слятия Chunked подход: единият изгражда цялата мрежа за сравнение и я задържа, а другият изгражда и изхвърля по един отрязък, така че паметта никога не нараства над размера на един отрязък.
Пуснахме с отворен код реализацията на Chunked загубата: github.com/CompactifAI/Full-Chunked-KL-Loss
Какво се променя на практика
Таблицата по-долу сравнява директно и четирите конфигурации: онлайн дестилацията и трите току-що описани реализации на offline загубата. При сравнение на един графичен процесор H200 с Llama 3.1 8B Instruct като учител и модел Llama 3.2B като ученик при контекст от 8K токена и четирите достигат почти идентична загуба при обучение, въпреки че offline изпълненията се обучават само спрямо кешираните top-100 логити за токен.
| Метод (контекст 8K, един H200) | Пикова памет | Време за итерация | Пропускателна способност |
|---|---|---|---|
| Онлайн дестилация | 102.8 GB | 25.9 s | 237 TFLOP/s |
| Offline, плътна KL | 78.3 GB | 18.5 s | 331 TFLOP/s |
| Offline, forward-chunked KL | 61.8 GB | 18.4 s | 335 TFLOP/s |
| Offline, слята Chunked KL | 58.3 GB | 20.2 s | 304 TFLOP/s |
Кривите на загубата се припокриват почти напълно и при четирите метода, потвърждавайки, че offline дестилацията с кеширани top-100 логити е без загуба спрямо онлайн дестилацията. Източник: Фигура 2 от статията. При тази дължина на последователността слятата Chunked загуба все още не е най-бързият вариант; допълнителната проекция при обратния проход коства малко скорост, но реалното ѝ предимство се проявява едва с нарастването на дължината на контекста, както показва следващият раздел.
Мащабиране към дълги контексти
За да видим по-ясно модела на мащабиране, проведохме изолиран тест върху мрежа за изходна проекция играчка (без тяло на трансформър, само ядрото на загубата). При 32K токена пиковата памет намалява от 85.2 GiB при плътната загуба до 5.45 GiB при напълно Chunked версията — намаление от 15,6 пъти — а плътната загуба се проваля изцяло при 64K токена и повече. При 256K токена напълно Chunked загубата използва 11.6 GiB спрямо 134.2 GiB при следващия най-добър Chunked вариант и е около 3,3 пъти по-бърза на итерация при тази дължина.
При дестилиране на модел GPT-OSS 20B в контекст от 32 768 токена освободената от слятата загуба памет позволи конфигурацията да бъде намалена от четири GPU възела до един. Времето за стъпка спадна от 57.0 до 12.23 секунди — около 5 пъти по-бързо — а пропускателната способност на GPU се увеличи от 74.2 до 345.7 TFLOP/s.
Полученият ученически модел
Именно ефективната offline конфигурация направи мащабната кампания за дестилация достъпна на първо място. Полученият компактен ученически модел, дестилиран от Llama 3.1 8B Instruct до около 3.2B параметъра, запазва по-голямата част от точността на учителя върху BoolQ и HellaSwag, а при MMLU остава на около девет пункта от него при по-малко от половината параметри.
Ученическият модел запазва по-голямата част от точността на учителя при кратък контекст и по-малко от половината размер. Източник: Фигура 6 от статията.
Multiverse Computing за практическо мащабно използване на дестилацията и възстановяването — не просто като еднократна рецепта, а като нещо, върху което екипите могат да работят итеративно на ниска цена. Статията обхваща и допълнителни аблации, например как изборът на функцията на загубата и пакетирането на последователностите влияят върху качеството на възстановяването.
Искате пълните технически подробности, включително градиента в затворен вид зад слятата Chunked загуба и цялата конфигурация за обучение? Прочетете пълната статия или се свържете с екипа ни, за да обсъдим прилагането на това във вашите собствени процеси за дестилация.
Също така пуснахме с отворен код реализацията на Chunked загубата: github.com/CompactifAI/Full-Chunked-KL-Loss
Модели, споменати в тази статия 6
Статии, споменати в тази статия 1
Колекции, споменати в тази статия 1
Общност
· Регистрирайте се или влезте, за да коментирате
Модели, споменати в тази статия 6
Статии, споменати в тази статия 1
Колекции, споменати в тази статия 1
Преведено автоматично от английски. Оригиналната статия е на връзката по-долу.

