在大詞彙語言模型中削減你的損失

Cut Your Losses in Large-Vocabulary Language Models

摘要

隨著語言模型變得越來越大,它們的詞彙表 (vocabulary) 也越來越大。這使得 LLM 訓練期間的記憶體佔用不成比例地移轉到單一一層:損失計算中的交叉熵 (cross-entropy)。交叉熵會建立一個 logit 矩陣,其中每一對「輸入 token × 詞彙項」都有一個項;對小模型而言,它消耗的記憶體比 LLM 其餘部分加起來還多一個數量級。我們提出 Cut Cross-Entropy (CCE),一種在不把所有 token 的 logit 實體化到全域記憶體 (global memory) 的情況下計算交叉熵損失的方法。相反地,CCE 只計算正確 token 的 logit、並即時 (on the fly) 對所有 logit 求 log-sum-exp。我們實作一個自訂核心 (kernel),在快閃記憶體 (flash memory) 中執行矩陣乘法與「對全詞彙的 log-sum-exp 歸約」,使交叉熵計算的全域記憶體消耗變得可忽略。這有戲劇性的效果。以 Gemma 2 (2B) 模型為例,CCE 把損失計算的記憶體佔用從 24 GB 降到 1 MB,並把分類器頭 (classifier head) 的訓練時總記憶體消耗從 28 GB 降到 1 GB。為改善 CCE 的吞吐量,我們利用 softmax 固有的稀疏性,提出跳過那些對梯度貢獻可忽略(即低於數值精度)的梯度計算元素。實驗證明,這種戲劇性的記憶體消耗降低,是在不犧牲訓練速度或收斂性的情況下達成的。 圖 1: 在 16-GPU(各 80 GB)完全分片資料平行 (fully-sharded data-parallel) 設定下、搭配激活檢查點 (activation checkpointing) 與混合精度 16-bit (fp16/bf16) AdamW 最佳化器,多種前沿模型的記憶體使用與可達最大批次大小(以百萬 token 計)。對每個模型,我們把其記憶體使用分解為「權重與最佳化器狀態」、「激活檢查點」、以及「交叉熵損失層計算的對數機率 (log-probabilities)」。我們的 Cut Cross-Entropy (CCE) 使批次大小能增加 1.5 倍(Llama 2 13B)到 10 倍(GPT-2、Gemma 2 2B),且不犧牲速度或收斂性。

此論文授權為 arXiv License

非 Creative Commons 授權的論文需登入後方可閱讀翻譯全文

Google 登入後閱讀

原文連結:arXiv