宅中地 - 每日更新
宅中地 - 每日更新

贊助商廣告

X

生物基礎模型的高效MoE訓練技術

2026年09月28日 首頁 » 熱門科技

隨著語言模型規模不斷擴大,密集架構的擴展成本變得越來越高昂。在密集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倍。

宅中地 - Facebook 分享 宅中地 - Twitter 分享 宅中地 - Whatsapp 分享 宅中地 - Line 分享
相關內容
Copyright ©2026 | 服務條款 | DMCA | 聯絡我們
宅中地 - 每日更新