annotated_deep_learning_paper_implementations 中的 MLP-Mixer:用序列混合 MLP 替换自注意力,训练 Masked Language Model 📅 发布时间:2026/9/5 19:45:45 👁 浏览次数: 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 ModelMLM 实验中的编码器自注意力完整实验代码见 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 层与可配置 TransformerMLPMixer 之所以能“几行代码”完成替换是因为仓库的 Transformer 实现把注意力模块抽象成了一个可注入参数。在 labml_nn/transformers/models.py 的TransformerLayer中采用 pre-norm 结构z self.norm_self_attn(x) self_attn self.self_attn(queryz, keyz, valuez, maskmask) x x self.dropout(self_attn)编码器每层以同一个张量作为 query/key/value 调用self_attn并以maskNone传入——这恰好满足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 Modelexperiment.py 在 MLM 实验 的基础上做最小改动把 MLP Mixer 接入训练流程。4.1 配置类继承 MLM 配置并新增混合 MLPclass 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_len32、mix_mlp.d_ff128计算单个混合 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 模型只使用编码器encodersrc_embedgenerator替换encoder_attn后整条前向链路就是字符嵌入 固定位置编码 → 若干层LayerNorm → 序列混合 MLP 残差 → LayerNorm → 逐位置 GELU FFN 残差→ 最终 LayerNorm → 线性层输出 logits逐层结构即 TransformerLayer。4.3 训练参数main()中的完整配置如下见 experiment.py 第 70–110 行配置项取值说明batch_size64每批 64 条长度seq_len的文本片段seq_len32序列长度取 32 以加快训练MLM 训练信号弱、周期长代码注释明确说明epochs1024训练 1024 个 epochinner_iterations1每 epoch 训练/验证切换 1 次d_model128token 嵌入维度transformer.ffn.d_ff256逐位置 FFN 隐藏层维度transformer.n_heads8头部数MLP 混合本身不使用多头从源码结构看MLM 模型只走编码器该值对混合层无实际影响transformer.n_layers6编码器层数transformer.ffn.activationGELU逐位置 FFN 激活函数mix_mlp.d_ff128序列混合 MLP 的隐藏层维度optimizer.optimizerNoam使用 Noam 优化器学习率按 step 衰减的调度方案optimizer.learning_rate1.0Noam 调度的基础学习率配置继承链为Configs→ MLM 的 Configs → NLPAutoRegressionConfigs → 训练/验证基础配置。因此除上表外还继承了 MLM 的默认设置masking_prob0.15随机掩蔽 15% 的 token、randomize_prob0.1其中 1/3 的掩蔽位置替换为随机 token、no_change_prob0.11/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),仅供参考