MOE Inference

FFN 理解

标准 FFN

$FFN(x) = max(0, xW_1 + b_1)W_2 + b_2$

$x$ 是 attention layer 的输出,第一步经过 $xW_1$ 将维度升高,对于升高后的维度使用 ReLU 对信息进行取舍,最后经过 $w_2$ 将信息降维回之前的维度

每一个 FFN 的过程就代表了一个专家,这个专家本质上包含了 $W_1$ 和 $W_2$ 两个参数矩阵

门控(Gated) FFN

$FFN(x) = \left( \text{Act}(x \cdot w_1) \odot (x \cdot w_3) \right) \cdot w_2$

简单来说和标准 FFN 一致,对升维后信息进行筛选,然后再降维。区别是筛选机制,标准直接采用的 ReLU,门控则是通过 gated

  • $x$: 上一层 Transformer/Self-Attention 模块输出的特征张量
  • $\odot$: Hadamard 积, 按元素逐个相乘
  • $w_1$: Gate Projection, 门控映射矩阵
  • $w_3$: Up Projection, 升维映射矩阵
  • $w_2$: Down Projection, 降维映射矩阵
  • $\text{Act}$: 非线性激活函数, 绝大多数情况下采用的是 SiLU(Sigmoid Linear Unit)

实际实现时,为了提高效率,会把 $w_1$ 和 $w_3$ 两部分结合起来计算。缺少的中间的Hadamard 积在实现上完成

$$[y_1, y_3] = x \cdot w_{13} = x \cdot [w_1, w_3]$$

从 deepseek-moe 模型看 MOE Architecture

  1. DeepseekModel 创建整个 transformer 架构,构建各个 block
1
2
3
self.layers = nn.ModuleList(
    [DeepseekDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
)
  1. DeepseekDecoderLayer 创建一个 block,包含 attention 和 moe layer
1
2
3
4
5
6
7
8
9
self.self_attn = Deepseek_ATTENTION_CLASSES[config._attn_implementation](config=config, layer_idx=layer_idx)

self.mlp = DeepseekMoE(config) if (config.n_routed_experts is not None and  \
    layer_idx >= config.first_k_dense_replace and layer_idx % config.moe_layer_freq == 0) \
    else DeepseekMLP(config)

self.input_layernorm = DeepseekRMSNorm(config.hidden_size, eps=config.rms_norm_eps)

self.post_attention_layernorm = DeepseekRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
  1. DeepseekMoE 创建一个 moe layer

transformer 骨架之外

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
class DeepseekForCausalLM(DeepseekPreTrainedModel):
    def __init__(self, config):
        self.model = DeepseekModel(config)
        self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)

    def forward():
            outputs = self.model(
                input_ids=input_ids,
                attention_mask=attention_mask,
                position_ids=position_ids,
                past_key_values=past_key_values,
                inputs_embeds=inputs_embeds,
                use_cache=use_cache,
                output_attentions=output_attentions,
                output_hidden_states=output_hidden_states,
                return_dict=return_dict,
            )
            logits = self.lm_head(hidden_states)

在 DeepseekModel 这个 transformer 骨架的基础上,还额外包含了 DeepseekForCausalLM 这样一层。从输出数据的内容来理解各自的区别会更简单一些

transformer 骨架生成的只是 hidden states,它的 shape 可能是 [batch, seq_len, hidden_size]

DeepseekForCausalLM 则通过 $logits = HW^{T}$,例如 [batch, seq_len, 2048] 的 H,[102400, 2048] 的 W,最终得到 [batch, seq_len, 102400] 的 logits。实际上这一步的目的是通过 lm_head 将 hidden state 从模型隐藏空间投影到 vocabulary space,也就是说这里的 102400 就是词表的大小

最后一步才是通过 softmax 获取到每个 token 的生成概率

0%