首页
/ annotated_deep_learning_paper_implementations 中的 MLP-Mixer:用序列混合 MLP 替换自注意力,训练 Masked Language Model

annotated_deep_learning_paper_implementations 中的 MLP-Mixer:用序列混合 MLP 替换自注意力,训练 Masked Language Model

2026-09-04 20:58:46作者:宗隆裙

本文基于仓库中 MLP-Mixer 模块 的文档与源码,讲解如何用十几行 PyTorch 代码实现论文《MLP-Mixer: An all-MLP Architecture for Vision》的核心思想:把 Transformer 的自注意力层替换为沿序列维度(token / 图像 patch 维度)作用的多层感知机(MLP)。读完后,你将理解 MLPMixer 模块如何通过转置张量实现对注意力层的“drop-in(即插即用)”替换,以及如何复用仓库的可配置 Transformer 与 MLM 训练框架,在 Tiny Shakespeare 语料上完整跑通一次实验。

1. MLP-Mixer 的核心思想

根据模块文档 readme.md 的说明,该模块是对论文 MLP-Mixer: An all-MLP Architecture for Vision 的 PyTorch 实现。论文将该模型应用于视觉任务:把输入图像切分为若干 patch,然后用施加在 patch 序列上的 MLP 替代注意力层——即整条网络完全由“特征混合 MLP + 序列(token)混合 MLP”两种组件堆叠而成,没有任何注意力机制。

文档中的关键结论是:

本仓库实现的 MLP Mixer 是 自注意力层 的 drop-in 替代品。它只是几行代码:把张量转置一下,使 MLP 沿序列维度而非特征维度作用。

虽然论文是在视觉任务上验证的,但本仓库将同样的模块搬到了自然语言方向——用它替代 Masked Language Model(MLM) 实验中的编码器自注意力,完整实验代码见 experiment.py

2. MLPMixer 模块的源码实现

核心实现全部位于 labml_nn/transformers/mlp_mixer/init.py,只有一个 MLPMixer 类:

class MLPMixer(nn.Module):
    def __init__(self, mlp: nn.Module):
        super().__init__()
        self.mlp = mlp

    def forward(self, query: torch.Tensor, key: torch.Tensor,
                value: torch.Tensor, mask: Optional[torch.Tensor] = None):
        # query, key, value 三者必须相同
        assert query is key and key is value
        # MLP mixer 不支持掩码
        assert mask is None

        x = query
        # 转置,使最后一维变成序列维度。
        # 新形状为 [d_model, batch_size, seq_len]
        x = x.transpose(0, 2)
        # 沿 token 维度施加 MLP
        x = self.mlp(x)
        # 转置回原形状
        x = x.transpose(0, 2)
        return x

从源码结构看,这里有三个值得注意的设计点:

  1. 刻意保留注意力接口forward(query, key, value, mask)MultiHeadAttention 的函数签名完全一致,输入形状同样为 [seq_len, batch_size, d_model]。这正是“drop-in replacement”的含义:上层调用者(Transformer 层)完全不需要感知被调用的是注意力还是 MLP 混合。源码还用 assert query is key and key is value 强制要求三者是同一对象(MLP 混合中 x = query = key = value),并用 assert mask is None 声明不支持任何掩码——所有 token 都能看到其他所有 token 的嵌入,这与双向 MLM 任务天然吻合,但也意味着该模块不能直接用于需要因果掩码的自回归解码。

  2. 两次转置实现“跨 token 的 MLP”nn.Linear 只作用于张量最后一维。输入 x 的形状是 [seq_len, batch_size, d_model],若直接送入 MLP,作用对象是 d_model(特征维度)——这正是普通 Transformer 中逐位置 FFN 的语义。x.transpose(0, 2) 之后形状变为 [d_model, batch_size, seq_len],最后一维变成序列长度,此时 self.mlp(x) 中的线性层权重在数学上就是作用于“所有位置”的矩阵,即 MLP-Mixer 论文中的 token-mixing MLP。计算完成后再转置回原形状。

  3. MLP 本身是外部注入的。构造函数只接收一个 nn.Module,模块不关心内部结构,实验中传入的是仓库通用的 位置前向网络 FeedForward(两层全连接 + 激活 + dropout)。

3. 接入点:Transformer 层与可配置 Transformer

MLPMixer 之所以能“几行代码”完成替换,是因为仓库的 Transformer 实现把注意力模块抽象成了一个可注入参数。在 labml_nn/transformers/models.pyTransformerLayer 中(采用 pre-norm 结构):

z = self.norm_self_attn(x)
self_attn = self.self_attn(query=z, key=z, value=z, mask=mask)
x = x + self.dropout(self_attn)

编码器每层以同一个张量作为 query/key/value 调用 self_attn,并以 mask=None 传入——这恰好满足 MLPMixer.forward 中的两条断言。

labml_nn/transformers/configs.py 中的 TransformerConfigs 更进一步,把注意力模块做成了可选项:encoder_attndecoder_attndecoder_mem_attn 默认值为 'mha'(对应 MultiHeadAttention 的计算函数 _mha)。'default' 选项下的 _encoder_layer 会用 c.encoder_attn 构造 TransformerLayer(见 configs.py 的 _encoder_layer),因此只需在实验配置中给 encoder_attn 赋一个 MLPMixer 实例,编码器各层就会自动装配 MLP 混合,其余部分(嵌入、逐位置 FFN、LayerNorm、堆叠逻辑)原封不动。

4. 完整实验:MLP Mixer + Masked Language Model

experiment.pyMLM 实验 的基础上做最小改动,把 MLP Mixer 接入训练流程。

4.1 配置类:继承 MLM 配置并新增混合 MLP

class Configs(MLMConfigs):
    # 可配置的位置前向网络,用作 MLP 混合层
    mix_mlp: FeedForwardConfigs


@option(Configs.mix_mlp)
def _mix_mlp_configs(c: Configs):
    """混合 MLP 的配置"""
    conf = FeedForwardConfigs()
    # 因为 MLP 是跨 token 施加的,
    # 所以 MLP 的“模型维度”设为序列长度
    conf.d_model = c.seq_len
    # 论文建议使用 GELU 激活
    conf.activation = 'GELU'
    return conf

注意 conf.d_model = c.seq_len 这一行:结合 FeedForward 的实现(layer1 = Linear(d_model, d_ff)layer2 = Linear(d_ff, d_model),线性层作用在最后一维),混合 MLP 实际是 Linear(seq_len -> d_ff) -> GELU -> Dropout -> Linear(d_ff -> seq_len)。以实验默认值 seq_len=32mix_mlp.d_ff=128 计算,单个混合 MLP 约 32×128 + 128×32 个权重,规模很小。

4.2 替换编码器注意力

@option(Configs.transformer)
def _transformer_configs(c: Configs):
    conf = TransformerConfigs()
    # 为嵌入与 logits 生成设置词表大小
    conf.n_src_vocab = c.n_tokens
    conf.n_tgt_vocab = c.n_tokens
    # 嵌入大小
    conf.d_model = c.d_model
    # 把注意力模块换成 MLPMixer
    from labml_nn.transformers.mlp_mixer import MLPMixer
    conf.encoder_attn = MLPMixer(c.mix_mlp.ffn)
    return conf

这里覆盖了父类 MLM 实验中的默认 _transformer_configs(默认使用 'mha')。由于 TransformerMLM 模型只使用编码器(encoder + src_embed + generator),替换 encoder_attn 后整条前向链路就是:字符嵌入 + 固定位置编码 → 若干层(LayerNorm → 序列混合 MLP 残差 → LayerNorm → 逐位置 GELU FFN 残差)→ 最终 LayerNorm → 线性层输出 logits,逐层结构即 TransformerLayer

4.3 训练参数

main() 中的完整配置如下(见 experiment.py 第 70–110 行):

配置项 取值 说明
batch_size 64 每批 64 条长度 seq_len 的文本片段
seq_len 32 序列长度取 32 以加快训练;MLM 训练信号弱、周期长,代码注释明确说明
epochs 1024 训练 1024 个 epoch
inner_iterations 1 每 epoch 训练/验证切换 1 次
d_model 128 token 嵌入维度
transformer.ffn.d_ff 256 逐位置 FFN 隐藏层维度
transformer.n_heads 8 头部数(MLP 混合本身不使用多头;从源码结构看,MLM 模型只走编码器,该值对混合层无实际影响)
transformer.n_layers 6 编码器层数
transformer.ffn.activation 'GELU' 逐位置 FFN 激活函数
mix_mlp.d_ff 128 序列混合 MLP 的隐藏层维度
optimizer.optimizer 'Noam' 使用 Noam 优化器(学习率按 step 衰减的调度方案)
optimizer.learning_rate 1.0 Noam 调度的基础学习率

配置继承链为:ConfigsMLM 的 ConfigsNLPAutoRegressionConfigs → 训练/验证基础配置。因此除上表外,还继承了 MLM 的默认设置:masking_prob=0.15(随机掩蔽 15% 的 token)、randomize_prob=0.1(其中 1/3 的掩蔽位置替换为随机 token)、no_change_prob=0.1(1/3 保持原 token 不变),掩蔽逻辑由 MLM 类 实现,损失只在被掩蔽的位置上计算([PAD] 位置被 CrossEntropyLoss(ignore_index=...) 忽略)。

4.4 运行方式

该实验基于仓库通用的 labml 实验框架(experiment.create / experiment.configs / experiment.start),安装依赖(见 requirements.txt)后,直接运行入口文件即可,实验会以 mlp_mixer_mlm 为名自动记录日志、定期采样生成文本并保存 PyTorch 模型:

python labml_nn/transformers/mlp_mixer/experiment.py

5. 小结与延伸阅读

这条从论文到代码的路径在仓库中非常清晰:

这套实现展示了该仓库的典型组织方式:把论文组件封装成与现有接口兼容的小模块,再借助 TransformerConfigs 的选项机制,用不到二十行实验代码完成“注意力 → MLP 混合”的架构替换,而数据管线、训练循环、采样与日志记录全部复用。需要留意其边界:MLPMixer 不支持掩码,因此只适合双向编码器场景(如这里的 MLM),不能用于需要因果掩码的自回归解码路径。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.12 K
2.72 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
904
1.82 K
docsdocs
暂无描述
Markdown
889
5.78 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
854
1.34 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
527
590
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.52 K
1.01 K
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.33 K
1.45 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
540
384
flutter_flutterflutter_flutter
本仓库是 Flutter SDK 与 Flutter Engine 的 OpenHarmony 适配版本,由 CPF-Flutter 团队维护。开发者可使用熟悉的 Flutter 技术栈开发 OpenHarmony 应用,3.35.7 及以后的适配版本可基于本仓库源码构建支持 OpenHarmony 的 Flutter Engine。
Dart
1.17 K
341