annotated_deep_learning_paper_implementations 中的 MLP-Mixer:用序列混合 MLP 替换自注意力,训练 Masked Language Model
【免费下载链接】annotated_deep_learning_paper_implementations🧑🏫 60+ Implementations/tutorials of deep learning papers with side-by-side notes 📝; including transformers (original, xl, switch, feedback, vit, ...), optimizers (adam, adabelief, sophia, ...), gans(cyclegan, stylegan2, ...), 🎮 reinforcement learning (ppo, dqn), capsnet, distillation, ... 🧠项目地址: https://gitcode.com/gh_mirrors/an/annotated_deep_learning_paper_implementations
本文基于仓库中 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从源码结构看,这里有三个值得注意的设计点:
刻意保留注意力接口。
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 任务天然吻合,但也意味着该模块不能直接用于需要因果掩码的自回归解码。两次转置实现“跨 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。计算完成后再转置回原形状。MLP 本身是外部注入的。构造函数只接收一个
nn.Module,模块不关心内部结构,实验中传入的是仓库通用的 位置前向网络 FeedForward(两层全连接 + 激活 + dropout)。
3. 接入点:Transformer 层与可配置 Transformer
MLPMixer 之所以能“几行代码”完成替换,是因为仓库的 Transformer 实现把注意力模块抽象成了一个可注入参数。在 labml_nn/transformers/models.py 的TransformerLayer中(采用 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_attn、decoder_attn、decoder_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.py 在 MLM 实验 的基础上做最小改动,把 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=32、mix_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 调度的基础学习率 |
配置继承链为:Configs→ MLM 的 Configs → NLPAutoRegressionConfigs → 训练/验证基础配置。因此除上表外,还继承了 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.py5. 小结与延伸阅读
这条从论文到代码的路径在仓库中非常清晰:
- 概念(“注意力换成跨 token 的 MLP”)→ labml_nn/transformers/mlp_mixer/readme.md;
- 核心模块(转置 + 注入 MLP + 两条断言)→ labml_nn/transformers/mlp_mixer/init.py;
- 被替换的参照物(多头注意力接口与实现)→ labml_nn/transformers/mha.py;
- 装配点(可配置 Transformer、pre-norm 层)→ labml_nn/transformers/configs.py、labml_nn/transformers/models.py;
- 任务侧(掩蔽策略与训练步)→ labml_nn/transformers/mlm/init.py、labml_nn/transformers/mlm/experiment.py;
- 完整可运行实验 → labml_nn/transformers/mlp_mixer/experiment.py。
这套实现展示了该仓库的典型组织方式:把论文组件封装成与现有接口兼容的小模块,再借助TransformerConfigs的选项机制,用不到二十行实验代码完成“注意力 → MLP 混合”的架构替换,而数据管线、训练循环、采样与日志记录全部复用。需要留意其边界:MLPMixer不支持掩码,因此只适合双向编码器场景(如这里的 MLM),不能用于需要因果掩码的自回归解码路径。
【免费下载链接】annotated_deep_learning_paper_implementations🧑🏫 60+ Implementations/tutorials of deep learning papers with side-by-side notes 📝; including transformers (original, xl, switch, feedback, vit, ...), optimizers (adam, adabelief, sophia, ...), gans(cyclegan, stylegan2, ...), 🎮 reinforcement learning (ppo, dqn), capsnet, distillation, ... 🧠项目地址: https://gitcode.com/gh_mirrors/an/annotated_deep_learning_paper_implementations
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考