- [object Object]Ascend 950PR/Ascend 950DT[object Object]:支持
- [object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:支持
- [object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]:支持
- [object Object]Atlas 200I/500 A2 推理产品[object Object]:不支持
- [object Object]Atlas 推理系列产品[object Object]:不支持
- [object Object]Atlas 训练系列产品[object Object]:不支持
接口功能:
Mega MoE算子 将 MoE 层的专家 FFN 的完整计算流程及前后数据通信(即 Dispatch + Linear1 + SwiGLU + Linear2 + Combine)融合为单个算子,实现了通信和计算的掩盖。 该算子提供了 mega_moe 与 get_symm_buffer_for_mega_moe等接口,这些接口需配套使用。
- get_symm_buffer_for_mega_moe:需与mega_moe配套使用,用于封装输入参数并创建SymmBuffer结构体,生成
[object Object]、[object Object]和[object Object]等mega_moe算子运行所需信息。
- get_symm_buffer_for_mega_moe:需与mega_moe配套使用,用于封装输入参数并创建SymmBuffer结构体,生成
计算公式:
输入:
- :激活矩阵,对应入参
[object Object]。 是全局总 token 数, 是隐藏层维度。 - :token 选择的专家编号矩阵,对应入参
[object Object]。 是每个 token 选择的专家数量。 - :token 选择的专家的门控权重矩阵,对应入参
[object Object]。 - :Linear1 的权重矩阵,对应入参
[object Object]。 是专家数量, 是中间层维度。 - :Linear2 的权重矩阵,对应入参
[object Object]。
- :激活矩阵,对应入参
输出:
- :最终输出矩阵,对应出参
[object Object]。
- :最终输出矩阵,对应出参
约定:
- 表示矩阵乘法, 表示逐元素乘法。
- 表示将 四舍五入到最近的整数, 表示将 向下取整。
- 表示取绝对值, 表示取最大值。
- 全体 token 的集合为 。
- 的 token 表示(即隐藏状态向量)为 ,且 。
- 的专家索引为 。
- 激活函数 ,其中 为 Sigmoid 函数。
- 。其中 的上标 表示对称量化值域区间:其值域关于 与 对称取整,与标准 INT8 的 值域不同,故以 上标区分。
- 张量切片操作采用Python风格的
[object Object]表示法,例如 代表取偶数行、 代表取奇数行。 - 表示二进制重解释操作,将张量的底层二进制数据按目标类型 重新解释。
计算过程
各产品支持的Linear计算时的激活矩阵A和权重矩阵W的数据类型如下:
[object Object]undefined
- [object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]、[object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:不支持上表中A8W8-FP场景。
- [object Object]Ascend 950PR/Ascend 950DT[object Object]:不支持上表中A16W16 、A8W8-INT、A8W4-INT场景。
EP Dispatch
在 Dispatch 阶段,每个 将其token表示 发送给专家 。即对于每个 ,专家 接收一份 。
记 为所有被分派给专家 的 token 索引集合,集合 的大小 即为专家 需要处理的 token 总数,则有 是由所有满足 的token表示 按任意固定顺序行堆叠而成的矩阵,该矩阵即为专家 经过 Dispatch 后收到的全部 token 表示。
对于每个 及其选中的第 个专家 ,存在唯一的行索引 ,使得 。该映射记录了 在专家 的输入矩阵 中的位置。
Expert Compute
在 MoE 层中,每个专家本质上是一个独立的前馈网络(FFN),采用 SwiGLU 结构以提升表达能力。整个计算过程分为如下三个子步骤。
1. Linear1 投影
Linear1 投影是专家网络的第一层线性变换,同时产生 gate 部分 和 up 部分 所需的预激活值。其计算公式为
2. SwiGLU 激活
首先将 沿列维度拆分为 gate 部分 和 up 部分 ,然后对 gate 部分应用 SiLU 激活函数,再与 up 部分逐元素相乘,得到专家 的中间激活表示 。
3. Linear2 投影
Linear2 投影作为第二层线性变换,将中间激活表示 从高维空间投影回原始的隐藏维度 ,使专家输出能够与残差连接等后续操作兼容。
经过以上计算,专家 的每一行输出对应其批次中的一个 token。对于 ,它在专家 中的输出行即为 。
Token Combine
Combine 负责收集所有专家计算出的输出向量,按照每个 token 原先分配到的专家权重进行加权求和,最终为每个 token 生成一个融合后的输出。利用之前记录的位置索引 ,从专家 的输出矩阵中收回属于 的行,并与门控权重相乘后求和:
其中 为 对专家 的门控权重。
所有 token 的输出按输入顺序堆叠为最终输出 。
输入
- :Linear1 权重矩阵的逐通道缩放因子,对应入参
[object Object]。 - :Linear2 权重矩阵的逐通道缩放因子,对应入参
[object Object]。
- :Linear1 权重矩阵的逐通道缩放因子,对应入参
EP Dispatch
在 Dispatch 通信之前,首先将原始 BF16 激活矩阵 量化为 INT8。对每个 ,计算其逐 token 缩放因子:
然后量化得到 INT8 表示:
在 Dispatch 通信阶段,每个 将其量化后的向量 和缩放因子 发送给专家 ()。
记 为所有被分派给专家 的 token 索引集合,集合 的大小 即为专家 需要处理的 token 总数。则有 是由所有满足 的 按任意固定顺序行堆叠而成的矩阵,该矩阵即为专家 经过 Dispatch 后收到的全部 token 表示;同理,对应的专家 收到的缩放因子向量记为 ,其元素由所有满足 的 按与 相同的行顺序堆叠而成。
对于每个 及其选中的第 个专家 ,存在唯一的行索引 ,使得 。该映射记录了 在专家 的输入矩阵中的位置。
Expert Compute
在 MoE 层中,每个专家本质上是一个独立的前馈网络(FFN),采用 SwiGLU 结构以提升表达能力。在A8W8场景下,两个线性层都使用 INT8 输入和 INT8 权重进行矩阵乘,得到 INT32 中间结果并反量化。具体分为三个子步骤。
1. Linear1 投影(INT8 矩阵乘 + 反量化)
Linear1 投影是专家网络的第一层线性变换,同时产生 gate 部分 和 up 部分 所需的预激活值。计算时执行 INT8 矩阵乘法,得到 INT32 计算结果:
然后反量化为预激活值 :
2. SwiGLU 激活
首先将 沿列维度拆分为 gate 部分 和 up 部分 ,然后对 gate 部分应用 SiLU 激活函数,再与 up 部分逐元素相乘,得到专家 的中间激活表示 。
3. Linear2 投影(量化 + INT8 矩阵乘 + 反量化)
Linear2 投影作为第二层线性变换,将中间激活表示 从高维空间投影回原始的隐藏维度 ,使专家输出能够与残差连接等后续操作兼容。 在 A8W8 场景下,需要将激活值量化为 INT8,因此先对 的每一行(每个 token)计算缩放因子:
得到专家 在 Linear2 计算时的激活缩放因子 。然后量化:
再执行 INT8 矩阵乘法并反量化:
经过以上计算,专家 的每一行输出对应其批次中的一个 token。对于 ,它在专家 中的输出行即为 。
Token Combine
Combine 负责收集所有专家计算出的输出向量,按照每个 token 原先分配到的专家权重进行加权求和,最终为每个 token 生成一个融合后的输出。利用之前记录的位置索引 ,从专家 的输出矩阵中收回属于 的行,并与门控权重相乘后求和:
其中 为 对专家 的门控权重。
所有 token 的输出按输入顺序堆叠为最终输出 。
输入
- :Linear1 权重矩阵的逐通道缩放因子,对应入参
[object Object]。 - :Linear2 权重矩阵的逐通道缩放因子,对应入参
[object Object]。 - :Linear1 的偏置矩阵,由 INT4 量化过程离线生成,对应入参
[object Object]。 - :Linear2 的偏置矩阵,由 INT4 量化过程离线生成,对应入参
[object Object]。
- :Linear1 权重矩阵的逐通道缩放因子,对应入参
EP Dispatch
在 Dispatch 通信之前,首先将原始 BF16 激活矩阵 量化为 INT8。对每个 ,计算其逐 token 缩放因子:
然后量化得到 INT8 表示:
在 Dispatch 通信阶段,每个 将其量化后的向量 和缩放因子 发送给专家
记 为所有被分派给专家 的 token 索引集合,集合 的大小 即为专家 需要处理的 token 总数。则有 是由所有满足 的 按任意固定顺序行堆叠而成的矩阵,该矩阵即为专家 经过 Dispatch 后收到的全部 token 表示;同理,对应的专家 收到的缩放因子向量记为 ,其元素由所有满足 的 按与 相同的行顺序堆叠而成。
对于每个 及其选中的第 个专家 ,存在唯一的行索引 ,使得 。该映射记录了 在专家 的输入矩阵中的位置。
Expert Compute
在 MoE 层中,每个专家本质上是一个独立的前馈网络(FFN),采用 SwiGLU 结构以提升表达能力。在A8W4-INT场景下,两个线性层都采用MSD(Mixed-precision Split-activation Decomposition,混合精度激活拆分分解)方案进行矩阵乘,通过将 INT8 值拆为高4位和低4位两个有符号 INT4,使得 INT8×INT4 的矩阵乘可分解为两个 INT4×INT4 的矩阵乘,从而利用硬件的 INT4 矩阵乘加速。该方案的数学原理和实现逻辑可参阅。
1. 生成精度补偿的偏置矩阵(离线生成,在算子外完成,并作为算子输入)
MSD 方案的核心是将量化后的 INT8 激活值二进制重解释地拆分为两个 INT4 分量。由于低位 INT4 分量按 定义,其数值范围被映射到 ,这相当于在原始无符号低 4 位(0~15)的基础上减去了 8。当将高、低位的 INT4 分别与权重矩阵做矩阵乘并合并时,该偏移会引入一个常数项,需要在最终结果中予以补偿。
记原始 INT8 激活矩阵为 ,拆分后的高位和低位 INT4 激活矩阵分别为 和 ,权重矩阵为 。恢复关系为:
其中 为与 形状相同的全 1 矩阵,用于逐元素加 8。令 为形状与输入特征维度一致的全 1 列向量,则矩阵乘展开为:
这里 表示对权重矩阵 的每一列求和,其结果为一个行向量(形状与输出维度相同)。在具体实现中,该行向量会被广播到批次维度,加到所有 token 的输出上。可见,若只计算前两项,会缺失一个正数项。为了精确恢复结果,需预先计算补偿偏置:
和 均按此法,分别使用对应的全 1 列向量(维度分别为 和 )与各自的权重矩阵相乘,得到形状匹配的偏置行向量。这些偏置在离线阶段计算完成后传入算子,在合并高低位结果时直接参与加法,从而消除偏移引入的误差。
2. Linear1 投影(INT4 × INT4 矩阵乘 + 反量化)
Linear1 投影是专家网络的第一层线性变换,同时产生 gate 部分 和 up 部分 所需的预激活值。
(1) 激活重解释为INT4
将 INT8 激活张量 二进制重解释为两个 INT4 交替拼接的视图:
其中高位和低位分量直接由 的偶数行和奇数行给出:
它们在数值上可以由原 INT8 经如下计算得到:
恢复关系为 。由于 仅改变类型视图, 与 共享底层物理内存,无需任何数据重排或拷贝。
(2) INT4 × INT4 矩阵乘与权重反量化
将重解释后的 INT4 激活视图 与权重矩阵 执行矩阵乘,并应用权重缩放因子 进行反量化,得到结果:
将结果按偶数行和奇数行分别记为如下的矩阵视图,它们即对应于和的计算结果:
(3) 精度补偿与激活反量化
利用偏置 和激活缩放因子 (行向量)分别进行精度补偿和激活反量化,最终预激活值为:
其中 表示将矩阵的每一行乘以 中对应的标量。
3. SwiGLU 激活
首先将 沿列维度拆分为 gate 部分 和 up 部分 ,然后对 gate 部分应用 SiLU 激活函数,再与 up 部分逐元素相乘,得到专家 的中间激活表示 。
4. Linear2 投影(量化 + INT4 × INT4 矩阵乘 + 反量化)
Linear2 投影作为第二层线性变换,将中间激活表示 从高维空间投影回原始的隐藏维度 ,使专家输出能够与残差连接等后续操作兼容。 在 A8W4 场景下,首先将激活值量化为 INT8,再通过二进制重解释为两个 INT4 的拼接视图,以一次矩阵乘完成计算,过程如下。
(1) 激活量化
对 的每一行计算缩放因子并量化至 INT8:
得到专家 在 Linear2 计算时的激活缩放因子 ,并计算量化结果:
(2) 激活重解释为 INT4
将 INT8 激活张量 二进制重解释为两个 INT4 交替拼接的视图:
其中高位和低位分量直接由 的偶数行和奇数行给出:
它们在数值上可以由原 INT8 经如下计算得到:
恢复关系为 。由于 仅改变类型视图, 与 共享底层物理内存,无需任何数据重排或拷贝。
(3) INT4 × INT4 矩阵乘与权重反量化
将重解释后的 INT4 激活视图 与权重矩阵 执行矩阵乘,并应用权重缩放因子 进行反量化,得到结果:
将结果按偶数行和奇数行分别记为如下的矩阵视图,它们即对应于 和 的计算结果:
(4) 精度补偿与激活反量化
利用偏置 和激活缩放因子 (行向量)分别进行精度补偿和激活反量化,最终输出为:
其中 表示将矩阵的每一行乘以 中对应的标量。
经过以上计算,专家 的每一行输出对应其批次中的一个 token。对于 ,它在专家 中的输出行即为 。
Token Combine
Combine 负责收集所有专家计算出的输出向量,按照每个 token 原先分配到的专家权重进行加权求和,最终为每个 token 生成一个融合后的输出。利用之前记录的位置索引 ,从专家 的输出矩阵中收回属于 的行,并与门控权重相乘后求和:
其中 为 对专家 的门控权重。
所有 token 的输出按输入顺序堆叠为最终输出 。
第一阶段对输入 Token 按专家分组收集后做 MXFP8 量化,生成各专家的量化输入与缩放因子:
说明:根据
[object Object]将 Token 按专家排序收集, 为分配到专家 的 Token 索引集合,表示当前专家收到的最大token数,每个专家数值可能不同, 为对应的子矩阵。 表示 MX 逐组量化(group size = 32),对每组 32 个元素提取共享指数后量化为 FP8 目标类型(FLOAT8_E5M2 或 FLOAT8_E4M3FN),同时输出 FLOAT8_E8M0 缩放因子。量化后的数据将作为 GMM1 的输入。第二阶段对每个专家执行 GMM1 矩阵乘法(将 沿列方向分为两半分别计算)、SwiGLU 激活和 MX 量化:
说明:将 的前 列 和后 列 分别与 MX 反量化后的输入做矩阵乘法,得到 Swish 分支 和门控分支 。SwiGLU 激活对两个分支做逐元素乘积 ,其中 为 Sigmoid 函数,将中间维度从 减半为 。随后对 SwiGLU 输出做 MX 量化,得到 GMM2 的量化输入 。
第三阶段对每个专家执行 GMM2 矩阵乘法,并将结果按目标 Rank 分发:
说明:将量化后的 SwiGLU 输出与第二组权重 做 MX 反量化后的矩阵乘法,将 维中间表示映射回 维隐藏空间,得到每个专家的输出 。计算完成后通过 RDMA peermem 将结果按目标 Rank 的专家偏移地址写入远端,实现跨 Rank 聚合。
第四阶段对所有 Token 按路由权重加权求和,恢复为与输入相同形状的输出:
说明:对每个 Token ,根据排序后的路由索引 从聚合后的专家结果中取出对应行,按
[object Object]中的权重逐元素加权累加,得到最终输出 。其中, 表示参数
[object Object], 表示参数[object Object], 表示参数[object Object], 表示参数[object Object], 表示参数[object Object], 表示属性[object Object](每个 Rank 的专家数), 表示[object Object]的第二维度(top-K 值,取值 6 或 8)。局部变量说明:
- :被路由到专家 的 Token 索引集合,由
[object Object]排序后确定。 - :专家 的量化输入及其 MX 缩放因子,第一阶段中间结果。
- 、: 对应专家 的前 列和后 列子矩阵,由
[object Object]按 SwiGLU 拆分推导。 - 、: 和 对应的 MX 缩放因子,从
[object Object]按维度截取。 - : 对应的 MX 缩放因子,来自参数
[object Object]。 - :GMM1 的两路矩阵乘法输出(Swish 分支和门控分支),中间结果。
- :SwiGLU 激活输出,维度 ,中间结果。
- :量化后的 SwiGLU 输出及其 MX 缩放因子,中间结果。
- :GMM2 的专家级输出,维度 ,中间结果。
- :Token 的第 个 top-k 专家在展开排序后的行索引,由路由排序确定。
- :MX 逐组量化操作,block size = 32,输出 FP8 数据和 E8M0 缩放因子。
- :MX 逐组反量化操作,在 matmul 内部隐式执行。[object Object]
- get_symm_buffer_for_mega_moe:
- mega_moe:
上角标[object Object]1[object Object]表示Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品不支持,上角标[object Object]2[object Object]表示Ascend 950PR/Ascend 950DT不支持,产品不支持的参数使用默认值即可。
[object Object]各张量参数的list[Tensor]长度、是否转置、是否支持非连续Tensor约束如下:
[object Object]参数一致性约束:
- mega_moe 接口的所有输入参数及其对应的张量维度,必须与 get_symm_buffer_for_mega_moe 的同名参数(例如
[object Object]、[object Object]、[object Object]等)保持一致。 - 调用算子过程中使用的
[object Object]、[object Object]、[object Object]、[object Object]、[object Object]等参数取值,所有卡需保持一致,网络中不同层中也需保持一致。
- mega_moe 接口的所有输入参数及其对应的张量维度,必须与 get_symm_buffer_for_mega_moe 的同名参数(例如
通信域和组网约束:
所有卡的
[object Object]、[object Object]参数取值需保持一致。各卡的通信域缓存区大小应当一致。
[object Object]为 HBM 上分配的 CCL 通信缓冲区总大小(Bytes),包含等大小的 windowIn 和 windowOut 两块空间,校验时以单个空间[object Object]为准,需满足:Atlas A2 训练系列产品/Atlas A2 推理系列产品:
[object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品:
[object Object]其中
[object Object]即通信域大小, 表示每张卡上可能专家数的最大值, 表示是否开启 dispatch 量化([object Object])。预留空间 10 MB 为内部元数据对齐与安全余量。通信域各节点的驱动版本应当相同。
Atlas A2 训练系列产品/Atlas A2 推理系列产品:多机通信域要求交换机组网,不支持双机直连组网。
Atlas A3 训练系列产品/Atlas A3 推理系列产品:多机通信域要求在一个超节点内,不支持双机直连组网和跨超节点组网。
参数约束:
- Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:
- 各卡
[object Object]需保持一致。 [object Object]:取值为[object Object]、[object Object]、[object Object]、[object Object]、[object Object]。[object Object]:取值范围为[object Object],且[object Object]。[object Object]:取值范围为[object Object],且[object Object]。[object Object]:取值范围为[object Object]。[object Object]需大于0,输入0表示自动计算,公式为[object Object]。[object Object]:取值范围为[object Object]。[object Object]:取值范围为[object Object],且[object Object]。[object Object]:取值范围为[object Object],且[object Object]。[object Object]:取值范围为[object Object](非量化)、[object Object](pertoken量化)。[object Object]:取值为[object Object]。- 支持三种计算场景(A16W16、A8W8-INT、A8W4-INT),不同场景下可选入参(缩放因子、偏置等)的必需性及数据类型有严格配套要求。调用时必须根据所选场景完整提供对应参数,不可混用或遗漏,配套关系见下表。
- 各卡
- Ascend 950PR/Ascend 950DT:
num_tokens(x.dim0)范围 [1, maxBs], maxBs = (8 (totalUbSize - 48 1024)) / (num_topk * (64 + ep_world_size)), 其中totalUbSize为UB总大小,950系列产品该值为256K。
hidden(x.dim1)仅支持1024、2048、3072、4096、5120、6144、7168、8192。
num_topk(topk_ids.dim1)支持[1, 16]。
num_experts_per_rank(weight1.dim0)范围 [1, 1024]。
intermediate_hidden(weight1.dim1)仅支持1024、2048、3072、4096、7168。
ep_world_size范围 [2, 1024]。
num_experts范围 [ep_world_size, 2048],且num_experts % ep_world_size == 0。
max_recv_token_num范围 [0, num_tokens × ep_world_size × min(num_topk, num_experts_per_rank)]。
dispatch_quant_out_dtype仅支持torch.float8_e5m2或torch.float8_e4m3fn或torch.float4_e2m1。
当前版本仅支持MXFP量化模式(dispatch_quant_mode = 4),dispatch阶段使用MX逐组量化(group size = 32),量化缩放因子类型为FLOAT8_E8M0。
x_active_mask和scales参数当前版本必须传入None,不支持非空输入。
combine_quant_mode当前支持0(非量化),3(MX模式float8_e5m2类型),4(MX模式float8_e4m3类型)。
comm_alg预留参数,必须为空字符串""。
y的数据类型与x相同。
weight1的dim1(intermediate_hidden)必须等于weight2的dim2的二倍,这是因为SwiGLU激活需要将中间维度从intermediate_hidden减半为intermediate_hidden/2。
weight_scales1和weight_scales2不可为空指针。
num_experts_per_rank = num_experts / ep_world_size,必须为整数且在 [1, 1024] 范围内。
weight_scales1和weight_scales2不可为空指针。
MXFP量化场景约束:
- weight1 shape为(num_experts_per_rank, intermediate_hidden, hidden),weight2 shape为(num_experts_per_rank, hidden, intermediate_hidden / 2)。
- weightScales1 shape为(num_experts_per_rank, intermediate_hidden, CeilDiv(hidden, 64), 2),其中 CeilDiv(hidden, 64) = ⌈hidden / 64⌉ = ⌊(hidden + 63) / 64⌋。
- weightScales2 shape为(num_experts_per_rank, hidden, CeilDiv(intermediate_hidden / 2, 64), 2),其中 CeilDiv(intermediate_hidden / 2, 64) = ⌈(intermediate_hidden / 2) / 64⌉ = ⌊(intermediate_hidden / 2 + 63) / 64⌋。
- weightScales1的dim3和weightScales2的dim3必须等于2。
- MXFP场景下,dispatch_quant_out_dtype=torch.float8_e5m2时weight1和weight2必须为FLOAT8_E5M2,dispatch_quant_out_dtype=torch.float8_e4m3fn时必须为FLOAT8_E4M3FN,dispatch_quant_out_dtype=torch.float4_e2m1时必须为FLOAT4_E2M1。
- x_active_mask和scales必须为None。
支持三种计算场景(A8W8-FP、A8W4-FP、A4W4-FP),不同场景下可选入参(缩放因子、偏置等)的必需性及数据类型有严格配套要求。调用时必须根据所选场景完整提供对应参数,不可混用或遗漏,配套关系见下表。
[object Object]
- Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:
默认支持确定性计算。
单算子模式调用:
下面示例将两个接口按调用顺序串联:先初始化通信域,再用 get_symm_buffer_for_mega_moe 构造 sym_buffer,最后调用 mega_moe 运行算子。
Ascend 950PR/Ascend 950DT:
[object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:
[object Object]