随着语言模型规模不断扩大,密集架构的扩展成本变得越来越高昂。在密集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文件中包含了启动、检查点保存和基准测试的相关命令。

欢迎在BioNeMo Recipes中尝试Mixtral原生Transformer Engine方案,并在NVIDIA Transformer Engine文档中了解更多关于优化MoE内核的信息。

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倍。

NVIDIA