资讯动态

在 annotated_deep_learning_paper_implementations 中实现 Pay Attention to MLPs(gMLP):门控 MLP 架构原理与自回归训练实战

发布时间:2026/9/18 10:03:28 来源:尧图企业网站定制
在 annotated_deep_learning_paper_implementations 中实现 Pay Attention to MLPsgMLP门控 MLP 架构原理与自回归训练实战【免费下载链接】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本文基于本仓库 gMLP 模块源码 与配套训练脚本 gMLP 实验完整讲解 Google 论文《Pay Attention to MLPs》arXiv:2105.08050提出的 gMLP 架构它用带门控gating的多层感知机替代 Transformer 的自注意力由 L 个 gMLP Block 堆叠而成。读完本文你将掌握 gMLP Block 与 Spatial Gating Unit 的逐行实现原理、权重初始化技巧、如何将 gMLP 无缝嵌入本仓库可配置的 Transformer 框架以及如何在 Tiny Shakespeare 数据集上启动一次端到端的字符级自回归训练。gMLP 是什么用门控 MLP 挑战自注意力论文《Pay Attention to MLPs》的核心论点是Transformer 的强大能力可能并不完全来自注意力机制本身而更多来自其感知机 信息混合的总体设计。基于这一观察论文提出了gMLPgated Multilayer Perceptron一个完全不含自注意力的纯 MLP 架构仅靠门控机制就能在序列维度上实现 token 间的信息交互并在语言建模等任务上达到与 Transformer 相当的效果。gMLP 模型由L 个 gMLP Block堆叠组成每个 Block 对输入嵌入 $X \in \mathbb{R}^{n \times d}$$n$ 为序列长度$d$ 为嵌入维度依次执行以下变换$$\begin{align} Z \sigma(XU) \ \tilde{Z} s(Z) \ Y \tilde{Z}V \end{align}$$其中 $U$、$V$ 为可学习投影权重$\sigma$ 为激活函数实现中使用 GeLU$s(\cdot)$ 即为下文详述的Spatial Gating Unit空间门控单元其输出维度是 $Z$ 的一半。在 本仓库的 gMLP 模块 中这一架构被实现为两个 PyTorch 模块GMLPBlockgMLP 块主体与SpacialGatingUnit空间门控单元并配套一个基于自回归语言建模任务的训练实验 experiment.py。GMLPBlock堆叠单元的实现GMLPBlock的构造函数接收三个关键参数见init.pyd_model输入嵌入维度 $d$d_ffn中间表示 $Z$ 的维度seq_lentoken 序列长度 $n$用于预分配空间门控的权重矩阵。块内部按以下顺序组织子模块normnn.LayerNorm([d_model])用于 Pre-Norm 残差结构先归一化再计算activationnn.GELU()即论文中的激活函数 $\sigma$proj1nn.Linear(d_model, d_ffn)实现 $Z \sigma(XU)$ 中的 $U$ 投影sguSpacialGatingUnit(d_ffn, seq_len)实现 $s(\cdot)$ 门控proj2nn.Linear(d_ffn // 2, d_model)实现 $Y \tilde{Z}V$ 中的 $V$ 投影输入维度为 $d_{ffn}/2$因为门控后通道数减半size固定为d_model用于与 Transformer 的Encoder模块对接详见第五节。前向过程forward(x, mask)init.py严格对应论文公式并带有残差连接shortcut x # 残差短路 x self.norm(x) # Pre-Norm 归一化 z self.activation(self.proj1(x)) # Z σ(XU) z self.sgu(z, mask) # Z̃ s(Z) z self.proj2(z) # Y Z̃V return z shortcut # 残差相加输入x的形状为[seq_len, batch_size, d_model]mask为形状[seq_len, seq_len, 1]的布尔掩码用于控制 token 之间的可见性在自回归场景下即只看过去、不看未来的下三角掩码。SpacialGatingUnitgMLP 的灵魂Spatial Gating UnitSGU是 gMLP 中唯一执行跨 token 信息交互的组件其数学形式为$$s(Z) Z_1 \odot f_{W,b}(Z_2)$$其中 $f_{W,b}(Z) WZ b$ 是沿序列维度的线性变换$\odot$ 为逐元素乘法$Z$ 沿通道嵌入维度被切分为等大的两部分 $Z_1$ 与 $Z_2$。本仓库的实现位于 SpacialGatingUnit核心细节如下。参数初始化接近恒等映射的起点论文特别强调$W$ 应初始化为小值、$b$ 初始化为 1这样在训练初期 $s(\cdot)$ 接近恒等映射除切分外保证训练稳定。源码严格遵循了这一点init.pyself.weight nn.Parameter( torch.zeros(seq_len, seq_len).uniform_(-0.01, 0.01), requires_gradTrue) self.bias nn.Parameter(torch.ones(seq_len), requires_gradTrue)即权重矩阵在 $[-0.01, 0.01]$ 内均匀初始化偏置恒为 1。此外$Z_2$在进入 $f_{W,b}(\cdot)$ 前还会经过一层nn.LayerNorm([d_z // 2])。前向计算与掩码机制前向过程init.py依次执行用torch.chunk(z, 2, dim-1)将 $Z$ 沿最后一维切分为 $Z_1$、$Z_2$若传入mask则做形状校验batch 维必须为 1即同一 batch 内所有样本共享掩码并去掉批维对 $Z_2$ 做 LayerNorm按当前实际序列长度截取权重子矩阵weight[:seq_len, :seq_len]从而支持输入序列长度不超过构造时seq_len的灵活调用若存在掩码将掩码与权重逐元素相乘——当 $W_{i,j}0$ 时$f_{W,b}(Z_2)$ 的第 $i$ 个位置就不会从第 $j$ 个 token 获取任何信息这正是因果约束的实现方式通过torch.einsum(ij,jbd-ibd, weight, z2) self.bias[:seq_len, None, None]完成 $W Z_2 b$返回 $Z_1 \odot (W Z_2 b)$。注意 $W$ 是逐 token 位置而非逐通道共享的矩阵这也是空间spatial一词的由来门控发生在序列方向。与 Transformer 框架的无缝集成gMLP 在本仓库中的一大工程亮点是无需重写模型骨架直接复用现有 Transformer 的Encoder、嵌入与生成头。实验代码 experiment.py 中的配置函数展示了这一点conf TransformerConfigs() # 见 labml_nn/transformers/configs.py conf.n_src_vocab c.n_tokens # 词表大小字符级 conf.n_tgt_vocab c.n_tokens conf.d_model c.d_model conf.encoder_layer c.gmlp # 关键用 gMLP Block 替换 Transformer Layer而TransformerConfigsconfigs.py中的encoder_layer默认构造TransformerLayer这里被整体替换为GMLPBlock。由于GMLPBlock提供了size属性且forward(x, mask)签名与TransformerLayer兼容它可以原样接入通用Encodermodels.py的层堆叠逻辑。配套地自回归模型AutoregressiveTransformerautoregressive_experiment.py负责通过src_embed带固定位置编码的嵌入层来自 models.py为输入加入位置信息生成下三角subsequent_maskutils.py确保每个 token 只能看到自身及之前的 token将输出送入generator线性层产生词表 logits。也就是说gMLP 的无注意力特性体现在块内部而模型的整体装配仍完整复用了 Transformer 的训练基础设施。自回归训练实验配置逐项解析训练脚本 experiment.py 继承了 basic/autoregressive_experiment.py 中的训练循环与配置基类最终继承自NLPAutoRegressionConfigs见 experiments/nlp_autoregression.py并为 gMLP 增加了d_ffn 2048这一专属配置项。其main()中的完整覆盖配置如下experiment.create(namegMLP) experiment.configs(conf, { # 数据与分词 tokenizer: character, # 字符级 tokenizer prompt_separator: , # 采样提示分隔符为空 prompt: It is , # 采样起始提示 text: tiny_shakespeare, # Tiny Shakespeare 数据集 # 序列与训练规模 seq_len: 256, # 上下文长度 256 epochs: 128, # 训练 128 个 epoch batch_size: 32, # 批大小 32 inner_iterations: 10, # 每个 epoch 内训练/验证切换 10 次 # 模型尺寸 d_model: 512, # 嵌入维度 d d_ffn: 2048, # gMLP 投影维度 d_ffn # 优化器 optimizer.optimizer: Noam, # Noam 学习率调度见 labml_nn/optimizers/noam.py optimizer.learning_rate: 1., # 基础学习率 })其中d_ffn通过GMLPBlock(c.d_model, c.d_ffn, c.seq_len)注入experiment.py。值得注意的几点字符级 tokenizer Tiny Shakespeare与本仓库 transformer 基础实验同源便于横向对比有注意力 vs 无注意力的建模能力差异Noam 优化器采用论文中常用的 warmup 衰减调度learning_rate1.是 Noam 调度的比例系数而非传统固定学习率Stochastic Depth论文中还使用了随机深度Stochastic Depth正则化训练时随机丢弃部分层实验注释明确说明本实现未包含该技巧——这一点在复现论文精度时需要留意。运行方式与扩展建议在安装好仓库依赖见 requirements.txt后直接执行即可启动训练python labml_nn/transformers/gmlp/experiment.py实验通过 labml 框架记录训练指标、保存模型检查点并支持在训练结束后用prompt如It is 进行采样生成。基于该实现可以从以下几个方向继续深入调整模型规模修改d_model、d_ffn、seq_len观察门控 MLP 在容量变化下的表现对照实验将encoder_layer换回默认的TransformerLayer删除conf.encoder_layer c.gmlp一行即可在相同数据与训练设置下与标准 Transformer 直接对比实现随机深度参考论文为Encoder的层堆叠加入随机丢弃补齐论文正则化细节探索掩码变体SpacialGatingUnit对掩码形状有明确断言batch 维为 1可自行扩展以支持批量内不同掩码。源码路径索引gMLP 核心实现GMLPBlock、SpacialGatingUnitlabml_nn/transformers/gmlp/init.pygMLP 自回归训练实验labml_nn/transformers/gmlp/experiment.py可配置 Transformer 框架labml_nn/transformers/configs.pyEncoder / 嵌入 / Generator 等通用组件labml_nn/transformers/models.py自回归训练基类labml_nn/transformers/basic/autoregressive_experiment.py下三角掩码工具labml_nn/transformers/utils.pyNoam 优化器labml_nn/optimizers/noam.py【免费下载链接】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),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价