隨著語言模型規模不斷擴大,密集架構的擴展成本變得越來越高昂。在密集Transformer中,每個Token都要經過每一層,因此增加能力就意味著訓練和推理的計算量都會隨之增加。
混合專家(MoE)架構採用了不同的擴展思路,通過使用大量子網路(即"專家"),但每個Token只激活其中一小部分專家。
這種權衡取捨使得MoE架構對大語言模型社區越來越具有吸引力。它們能夠更高效地擴展模型容量,但收益很大程度上取決於具體實現方式。碎片化的專家計算會降低GPU利用率,路由機制會增加通信開銷,而更大的參數規模則會帶來內存和分布式訓練方面的挑戰。NVIDIA Transformer Engine(TE)通過針對分組專家計算、核心融合和低精度訓練優化的原語,幫助解決這些瓶頸問題。隨著生物基礎模型的參數量和序列長度不斷增長,這些原語能夠在擴展模型容量的同時提升GPU效率。
本教學展示了如何藉助NVIDIA BioNeMo MoE方案和TE將這些技術付諸實踐。您將看到GroupedLinear如何改進專家計算、MXFP8如何降低內存占用,以及GroupedMLP核心如何將量化、SwiGLU和路由權重縮放融合在一起。這些能力共同為高效訓練基於MoE的生物基礎模型提供了實用參考。
前提條件
在開始之前,您需要具備以下條件:
熟悉Python、PyTorch和分布式訓練概念
一個支持NVIDIA CUDA的環境——您可以使用鏈接中提供的Dockerfile,或自行安裝該方案所需的依賴項
至少兩塊GPU用於專家並行;若要使用融合的MXFP8 GroupedMLP核心,則需要NVIDIA Blackwell GPU
挑戰一:碎片化的專家核心
MoE模型用多個專家網路替代了單一的密集前饋模組。然而,簡單粗糙的實現方式可能會引發過多的核心啟動。例如,Hugging Face的基準實現是在Python循環中遍歷所有專家,每個專家都會觸發單獨的核心啟動。
```
for expert_idx, expert_layer in enumerate(self.experts):
idx, top_x = torch.where(expert_mask[expert_idx])
current_state = hidden_states[None, top_x].reshape(-1, hidden_dim)
current_hidden = expert_layer(current_state) * routing_weights[top_x, idx, None]
final_hidden_states.index_add_(0, top_x, current_hidden)
```
分組執行保留了各個專家矩陣的獨立性,但將它們的計算工作合併提交。TE的GroupedLinear通過匯集專家權重和輸入Token,在一次調用中完成多個線性變換。由於每個專家接收的Token數量可能不同,GroupedLinear接受按專家劃分的Token數量參數(split_sizes)。它通過TE的分組GEMM路徑提交本地專家計算,而不是為每個專家單獨啟動一次PyTorch線性運算,從而減少了啟動和調度開銷。
GroupedLinear的使用方式如下。每個專家保留自己的權重張量(weight0、weight1等),調用時將按專家劃分的Token數量作為額外的位置參數傳入:
```
from transformer_engine.pytorch.ops import GroupedLinear
experts_gate_up = GroupedLinear(
num_groups=num_local_experts,
in_features=hidden_size,
out_features=2 * intermediate_size,
bias=False,
dtype=torch.bfloat16,
device="cuda",
)
gate_up_output = experts_gate_up(tokens, split_sizes)
```
與Python循環相比,這種方式將門控-升維投影作為一次分組操作提交,而不是多次單獨調用。
Hugging Face Transformers也提供了grouped_mm功能。不過,如後文所述,TE可以將GroupedLinear與MXFP8量化、激活函數、路由權重縮放以及中間數據搬移融合為一個GroupedMLP核心。
挑戰二:模型規模龐大與激活內存占用
MoE架構增加了總參數容量,而基因組學工作負載通常使用長序列,這給訓練過程中的激活內存帶來了壓力。BF16使用16位來表示每個模型權重和激活值。
BioNeMo方案藉助TE支持FP8和MXFP8訓練,以降低內存占用。這兩種格式都使用8位而非16位來表示權重和激活值。FP8與MXFP8的主要區別在於縮放粒度:MXFP8為每32個連續數值分配一個縮放因子,有助於保持數值範圍和精度。在NVIDIA Blackwell GPU上,MXFP8獲得了硬體加速支持,使MXFP8的GEMM運算能夠使用專門的Tensor Core指令。有關MXFP8和分塊縮放的詳細資訊,請參閱Transformer Engine FP8入門指南。
挑戰三:低精度訓練中的量化開銷
儘管大部分訓練計算使用8位精度,但模型仍以16位保留其主權重。因此,訓練框架需要增加量化和反量化步驟以在不同格式間進行轉換。量化會在低精度GEMM運算之前將BF16權重和激活值轉換為MXFP8;反量化則將結果轉換回更高精度的格式。簡單粗糙的實現方式會將這些步驟作為獨立操作執行,這也正是後文所述融合MLP路徑的設計動機。
```
fp8_recipe = te_recipe.MXFP8BlockScaling()
model = TEMixtralMXFP8ForCausalLM(config, fp8_recipe=fp8_recipe, dispatcher=dispatcher)
```
TE的autocast API可以為模型的前向和反向傳播啟用MXFP8精度:
```
with te.autocast(enabled=True, recipe=self._fp8_recipe):
for decoder_layer in self.layers:
hidden_states = decoder_layer(hidden_states)
```
完整代碼請參閱BioNeMo方案。
要使用融合MLP,需導入Transformer Engine的Sequential API,將gate_up、ScaledSwiGLU和down串聯在一起。該API還會將反量化步驟摺疊進融合路徑中。ScaledSwiGLU將路由概率("縮放因子")與專家前饋網路計算結合在一起。
```
from transformer_engine.pytorch.ops import GroupedLinear, ScaledSwiGLU, Sequential
experts_ffn = Sequential(GroupedLinear(gate_up), ScaledSwiGLU(), GroupedLinear(down))
```
TE的Sequential API會掃描這些操作,當模式匹配時,將GroupedLinear→ScaledSwiGLU→GroupedLinear這一序列替換為一個融合操作對象:前向傳播對應ForwardGroupedMLP_CuTeGEMMSwiGLU_MXFP8,反向傳播則對應相匹配的融合反向操作。這減少了框架開銷,將SwiGLU和概率縮放工作融合進分組MLP路徑中,並避免了部分中間結果的具體化生成。
成果
以上是BioNeMo方案中的部分優化措施。在我們基於八塊NVIDIA B200 Tensor Core GPU進行的訓練基準測試中,該方案實現的吞吐量最高可達Hugging Face基準的2.21倍。
運行該方案
首先使用雙GPU的L0_sanity配置,確認專家並行和訓練環境運行正常:
```
torchrun --nproc_per_node=2 train_fsdp2_ep.py --config-name L0_sanity
```
驗證完成後,可擴展至Mixtral-8x7B配置,在八塊GPU上採用專家並行(EP=8)和MXFP8精度:
```
torchrun --nproc_per_node=8 train_fsdp2_ep.py --config-name L1_8x7B_ep checkpoint.ckpt_dir=/path/to/ckpt
```
根據您的GPU和內存需求選擇BF16或MXFP8,並設置數據並行和專家並行的規模,使二者乘積等於GPU總數。該方案的README文件中包含了啟動、檢查點保存和基準測試的相關命令。
Q&A
Q1:什麼是MoE(混合專家)架構?它有什麼優勢?
A:MoE是一種模型擴展架構,使用多個專家子網路,但每個Token只激活其中一小部分。相比密集架構,它能更高效地擴展模型容量,但需要良好的實現方式才能真正發揮優勢。
Q2:GroupedLinear是如何提升專家計算效率的?
A:GroupedLinear通過匯集專家權重和輸入Token,在一次調用中完成多個線性變換,而不是為每個專家單獨啟動核心。這大幅減少了核心啟動和調度開銷,相比傳統的Python循環遍歷方式效率更高。
Q3:MXFP8精度訓練能帶來什麼好處?
A:MXFP8使用8位而非16位表示權重和激活值,能顯著降低內存占用。在NVIDIA Blackwell GPU上,MXFP8還獲得了硬體加速,配合融合核心技術,訓練吞吐量最高可提升至基準的2.21倍。






