annotated_deep_learning_paper_implementations 中的 MLP-Mixer:用序列混合 MLP 替换自注意力,训练 Masked Language Model
本文基于仓库中 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.py
5. 小结与延伸阅读
这条从论文到代码的路径在仓库中非常清晰:
- 概念(“注意力换成跨 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),不能用于需要因果掩码的自回归解码路径。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0622
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00