Day 3 我用「每參數 12 bytes」估出 4B 級模型可以在 128GB 上做全參數訓練。今天要把估算換成實測,因為估算永遠會少算一塊:激活值。
DGX Spark(GB10)的關鍵特性是統一記憶體架構:CPU 與 GPU 共享同一池記憶體,沒有 PCIe 傳輸瓶頸,但也代表你不能用「把東西丟到 CPU RAM」來省 GPU 記憶體——那是同一池。
軟體端我用的組合是標準的:PyTorch、Hugging Face Transformers、TRL / Accelerate、DeepSpeed(主要用 ZeRO stage 的梯度分片邏輯,即使單卡也有意義)。
我建議第一件事就是把環境固定下來——把套件版本鎖成一份可重現的清單。理由不是潔癖,是你的訓練會跑好幾天,中間如果因為任何原因要重跑,你需要環境完全一致,否則你分不清結果差異來自資料還是來自套件更新。
bf16 下,每參數 2 bytes。這塊是固定的,沒有調整空間(除非量化,但量化訓練是另一個題目)。
同樣 bf16,每參數 2 bytes。也是固定的。
這塊是最大的一塊,也是最有調整空間的一塊。
AdamW 需要保存一階動量與二階動量。如果用 fp32 保存,每參數 8 bytes;再加上如果保存 fp32 的主權重副本,再多 4 bytes。
可以動的選項:
我最後的選擇是保留 AdamW 但用 8-bit 版本,因為我要三階段都跑得動,而 CPT 階段最吃緊。
激活值的佔用跟 batch_size × sequence_length × hidden_size × 層數 成正比。這是唯一一塊跟資料形狀有關的記憶體,也是最容易 OOM 的地方。
三個控制手段:
坑一:一開始沒開 gradient checkpointing,直接 OOM。
這個很好解,但我要講的是它的診斷過程——OOM 的錯誤訊息會告訴你「試圖分配 X GB 但只剩 Y GB」,而 X 這個數字如果遠大於你的權重大小,那八成就是激活值。
坑二:sequence length 設太長,看似沒 OOM,但速度慢到不可接受。
我一開始想「反正記憶體夠,設長一點比較好」。結果是每步訓練時間拉長好幾倍,整個 CPT 從預估兩天變成一週。
我的處理:先做一次資料的長度分佈統計。大部分資安文本落在哪個區間?如果 90% 的樣本都在某個長度以下,那把上限設在那裡,剩下 10% 做切分,遠比為了少數長文拖累全部來得划算。
坑三:忘了留空間給評測。
訓練跑滿記憶體之後,我想同時跑一下評測——OOM。後來的做法是評測完全獨立跑,訓練結束後才跑,不要貪快。
在正式開跑之前,我用 1% 的資料跑完整個三階段。
跑完的結果當然沒有意義(模型什麼都沒學到),但它驗證了:
這一步花了幾小時,但它避免的是「跑到第 40 小時才發現第三階段的 checkpoint 格式讀不進去」這種事。
明天正式開始 CPT。
🛡️ Instagram: @aid3fend — AI 資安實戰紀錄,歡迎追蹤交流。