知識蒸餾是一種機器學習領域的知名技術,透過訓練一個較小的學生模型來匹配較大教師模型的效能。隨著近期開源大型語言模型(LLM),如 gpt-oss、Qwen、GLM 或 Kimi 的興起,知識蒸餾再次成為主流研究主題。

部署這些超大型模型成本高昂:最近的 Kimi-K3 模型擁有 2.8 兆個參數,僅載入就需要約 3TB 的顯示記憶體(VRAM)。因此,將其壓縮成較小的模型並透過知識蒸餾恢復原始能力,已成為標準做法,Nvidia (Nemotron 3 Puzzle 75B) 或 Multiverse Computing (Hypernova 60B) 等公司近期也發布了高品質的壓縮模型。

蒸餾步驟決定了最終模型的大部分品質,但它通常也是整個流程中最昂貴的部分。同時載入教師和學生模型,並為每個 token 產生整個詞彙表的機率分佈,需要大量的 VRAM,通常只有數百個 GPU 和精密的張量平行策略才能實現。

我們最新的論文《LLM 的高效知識蒸餾:離線 Top-K Logits 和融合分塊 KL 損失》透過兩項系統變革解決了這個問題:一次性快取教師模型的 Top-K logits,這樣教師模型就不必與學生模型同時存在於記憶體中;以及一種新的、記憶體高效的 KL 散度損失,它避免了實體化完整的「詞彙量 × 序列長度」矩陣。

這兩項改變將 VRAM 使用量大幅降低,遠超 PyTorch 或 NVIDIA Megatron-Bridge 等函式庫的預設實作。

這兩項改變共同將訓練成本降低到足以讓單一 GPU 上實現長上下文修復,並使其足夠便宜,讓大規模實驗變得實用。

標準的設定,即使用 Kullback-Leibler 散度損失(KL 損失)的線上蒸餾,需要同時載入教師和學生模型。在每個訓練步驟中,教師模型會執行完整的前向傳播以產生其輸出分佈,而學生模型則被訓練來匹配它。

這是最具表達力的設定,因為可以獲得完整的教師分佈,但它也是記憶體和計算最密集的:每個 token 位置必須保留兩個完整的詞彙表張量,而且教師模型在每次訓練步驟中都必須重新計算,儘管其行為在整個訓練過程中並未改變。

舉一個實際例子,gpt-oss-120b 擁有 201,088 個 token 的詞彙表。在序列長度為 32K 且批次大小為 4 的情況下,僅教師機率張量就有 4 × 201,088 × 32,768 的形狀;以 bfloat16 格式計算,單一張量就需要約 50GB 的 VRAM。

如果加上梯度、激活、模型權重和優化器狀態,單次蒸餾訓練迭代的 VRAM 峰值可能達到約 250GB,這甚至超過了 H200 或 B200 GPU 所能提供的容量。在這篇文章中,我們展示了透過重新設計 KL 損失以分塊處理數據,可以將此成本降至幾乎為零。

密集型 KL 損失的 VRAM 峰值約為 250GB,超過了單一 H200 GPU 的 141GB 容量。而融合分塊損失則從未產生這種峰值,最高約為 128GB。來源:論文圖 1。

我們提出了兩項系統變革。

首先是離線蒸餾。我們不再在每個步驟中重新計算教師模型,而是將其輸出計算一次,快取每個位置最有可能的 Top-100 token,然後讓學生模型針對該快取進行訓練。教師模型在訓練期間無需駐留在記憶體中,一旦快取建立,也無需再次運行,因此相同的快取可以重複用於多個消融實驗。

其次是融合分塊 KL 損失。要了解損失本身為何昂貴,想像一下它實際構建的內容:對於序列中的每個 token 位置和詞彙表中的每個單詞,損失需要一個數字來描述學生模型的預測與教師模型的分歧程度。如果將其佈局為網格,那就是每個詞彙條目一行,每個序列位置一列。對於一個擁有 10 萬以上單詞的詞彙表和一個長序列來說,這個網格是巨大的,而計算 KL 損失的預設方式是在產生單一數字之前構建整個網格。

我們比較了三種計算相同損失的方法,它們在數學上是等效的:

密集型 KL 損失是教科書上的方法。它從快取的 Top-100 logits 重建一個完整的、密集的教師機率網格,並將其與學生模型自身的密集型 log-機率網格進行比較。這是最接近線上蒸餾工作方式的版本,因此我們將其作為正確性基準,但它將完整的「詞彙量 × 序列」網格在記憶體中保留了兩次。

前向分塊 KL 損失保持教師模型稀疏(每個位置僅快取其 Top-100 logits,從不擴展為密集網格),並逐塊計算損失,一次處理一個序列位置切片。這消除了密集的教師模型和密集的比較,並且在我們的基準測試中證明是三種方法中最快的。然而,它仍然有一個盲點:學生模型自身的 logits,即模型輸出層產生的網格,仍然會完整計算並保留用於反向傳播,因此記憶體仍會隨著序列長度急劇增長。

融合分塊 KL 損失是我們的主要貢獻,它更進一步,將模型的輸出投影直接融合到損失計算中。它根本不產生學生模型的完整 logits 網格:它一次處理序列的一個分塊,將隱藏狀態投影到該分塊的 logits,將結果納入運行中的損失,然後在移動到下一個分塊之前丟棄該分塊。

反向傳播會即時重新計算每個分塊,而不是儲存它。代價是進行兩次投影,一次前向,一次反向,但作為交換,峰值記憶體僅隨序列長度線性增長,而不是隨著完整的「詞彙量 × 序列大小」而飆升。

下面的 GIF 顯示了密集型和融合分塊方法的區別:一個構建整個比較網格並保留所有內容,另一個則一次構建和丟棄一個切片,因此記憶體從未超出單個分塊。我們已將分塊損失的實作開源:github.com/CompactifAI/Full-Chunked-KL-Loss。

這在實踐中帶來了哪些改變?

下表將所有四種設定進行了比較:線上蒸餾,以及剛剛描述的三種離線損失實作。在單一 H200 GPU 上,以 Llama 3.1 8B Instruct 作為教師模型,3.2B Llama 模型作為學生模型,在 8K token 上下文長度下進行比較,所有四種方法都達到了幾乎相同的訓練損失,儘管離線運行僅針對每個 token 快取的 Top-100 logits 進行訓練。

方法 (8K 上下文, 單一 H200) | 峰值記憶體 | 迭代時間 | 吞吐量

--|--|--|--

線上蒸餾 | 102.8 GB | 25.9 s | 237 TFLOP/s

離線, 密集型 KL | 78.3 GB | 18.5 s | 331 TFLOP/s

離線, 前向分塊 KL | 61.8 GB | 18.4 s | 335 TFLOP/s

離線, 融合分塊 KL | 58.3 GB | 20.2 s | 304 TFLOP/s

損失曲線在所有四種方法中幾乎完全重疊,證實了使用 Top-100 快取 logits 的離線蒸餾相對於線上蒸餾是無損的。來源:論文圖 2。在這種序列長度下,融合分塊損失還不是最快的選項,其額外的反向傳播投影會稍微犧牲速度,但其真正的優勢只有在上下文長度增加時才會顯現,下一節將會展示。

為了更清晰地看到擴展模式,我們在一個玩具輸出投影網路(沒有 Transformer 主體,只有損失核心)上進行了獨立基準測試。在 32K token 時,峰值記憶體從密集型損失的 85.2 GiB 降至完全分塊版本的 5.45 GiB,減少了 15.6 倍,而密集型損失從 64K token 開始就完全失敗了。

在 256K token 時,完全分塊損失使用了 11.6 GiB,而次優的分塊變體則使用了 134.2 GiB,並且在該長度下每次迭代速度約快 3.3 倍。

在 32,768 token 上下文下蒸餾 GPT-OSS 20B 模型時,融合損失釋放的記憶體使設定從四個 GPU 節點縮減到一個。步驟時間從 57.0 秒降至 12.23 秒,約快了 5 倍,每個 GPU 的吞吐量從 74.2 TFLOP/s 提升到 345.7 TFLOP/s。

最終的學生模型表現如何?

高效的離線設定是使大規模蒸餾活動變得經濟實惠的根本原因。由此產生的緊湊型學生模型,從 Llama 3.1 8B Instruct 蒸餾到約 3.2B 參數,在 BoolQ 和 HellaSwag 上保留了教師模型的大部分準確性,在 MMLU 上與教師模型相差約九個點,而參數數量不到一半。

學生模型在不到一半的尺寸下,保留了教師模型大部分的短上下文準確性。來源:論文圖 6。

這項工作是 Multiverse Computing 正在進行的研究的一部分,旨在使蒸餾和修復在規模化運行時變得實用,不僅僅是一次性的方法,而是團隊可以廉價地迭代的東西。該論文還涵蓋了額外的消融實驗,例如損失函數的選擇和序列打包如何影響恢復品質。

想要完整的技術細節,包括融合分塊損失背後的閉合形式梯度和完整的訓練配置嗎?請閱讀完整論文,或聯繫我們的團隊討論如何將此應用於您自己的蒸餾流程。我們也已將分塊損失的實作開源:github.com/CompactifAI/Full-Chunked-KL-Loss。