TriPLU:在微型语言模型中通过直接三线性积前馈网络绕过门控

arXiv cs.CL 论文

摘要

TriPLU 为微型仅解码器语言模型引入了直接三线性积前馈网络,展示了在低计算环境下验证损失优于 SwiGLU 的结果。

arXiv:2608.20360v1 Announce Type: new 摘要:我们研究微型仅解码器语言模型是否从直接乘以学习特征投影的前馈层中获益。TriPLU,一种三线性积线性单元,用仅乘积的度-3分支替代通常的门控FFN分支,该分支逐坐标乘以三个投影流。在字符级 TinyStories 1M字节前缀研究中,TriPLU 达到平均最佳验证损失 1.0637,相比之下,匹配的 SwiGLU 为 1.1017,度-4 乘积控制为 1.0780,度-2 控制为 1.1026。在仅训练的 Byte-BPE 实验中,TriPLU 在低学习率设置下也降低了 TinyStories 和 WikiText-2 原始数据上的验证和保留字节位,PMI切片证据表明在已见的中、高PMI相邻词对上有所增益。恒定学习率诊断表明,乘积分支归一化可以减少高学习率最佳检查点差距,尽管在热调度下最终BPB仍然下降。因此,结论是刻意狭隘的:直接乘积FFN可以在特定低计算环境下改善固定预算小模型的损失,但该分支对优化敏感,并未建立FLOP归一化效率、缩放行为或广泛的LLM性能。
查看原文
查看缓存全文

缓存时间: 2026/08/24 04:11

# TriPLU:通过直接三线性乘积前馈网络绕过门控机制于微型语言模型中
来源:https://arxiv.org/html/2608.20360
###### 摘要
我们探究微型仅解码器语言模型是否受益于能直接对学习到的特征投影进行乘法的前馈层。TriPLU(三线性乘积线性单元)是一个纯乘积的三阶前馈网络分支,它对三个投影流进行逐坐标乘法。在一项基于字符级的TinyStories 1M字节前缀研究中,TriPLU达到了平均最佳验证损失1.0637,而参数匹配的SwiGLU为1.1017,四阶乘积控制组为1.0780,二阶控制组为1.1026。在仅训练集的Byte-BPE实验中,TriPLU在低学习率设置下也降低了TinyStories和WikiText-2原始文本上的验证集和保留集每字节比特数,其PMI切片证据与在已见的中等和高PMI相邻词对上的增益相符。恒定学习率诊断显示,乘积分支归一化可以减少高学习率最佳检查点差距,但在热调度下最终每字节比特数仍会下降。其主张范围有限:直接乘积前馈网络在特定的低算力配置下可以改善固定预算的小模型损失,但该分支对优化敏感,且未确立FLOP归一化效率、扩展性或广泛的LLM性能。††AI使用声明:生成式人工智能被用于编码支持、文献检索协助、实验日志组织及文本修订。人类作者审阅并编辑了手稿,核实了来源与证据,并对所有方法、结果、主张、代码和写作负责。

## 1 引言
Transformer前馈网络通常结合仿射投影与逐元素非线性激活。它们可以通过组合来近似乘法特征交互,但小模型可能需要额外的宽度、深度或数据来学习乘积偏向层直接表达的交互。现代Transformer前馈网络在门控变体中已包含乘法结构。Shazeer (2020) 表明,包括GEGLU和SwiGLU在内的GLU变体,在序列到序列任务中相比ReLU或GELU能改进Transformer前馈子层。这使得SwiGLU成为一个必要的基准,而非边缘对比。本文的问题更具体:当参数和训练词元数匹配时,一个对学习投影更显式的乘积分支,在微型仅解码器语言模型中,是否能比一个强大的门控前馈网络带来额外增益?“绕过门控”意味着用直接乘积路径替换前馈网络通常的激活与门控形式。TriPLU不对一个流应用一元激活来门控另一个流。相反,它直接将三个学习到的投影相乘,并让这个三线性乘积提供非线性。显式乘积为特征共现项提供了短路径。我们探究直接乘积前馈网络是否在微型仅解码器语言模型中改善验证和保留损失,对抗紧密匹配的SwiGLU基准和乘积阶控制,以及增益是否出现在乘积偏向应起作用的共现上下文中。乘法神经单元并非新事物:门控、注意力、超网络、动态层和神经算术模块都使用了相关思想。我们的贡献是一项受控的低算力建模研究,而非一个新的算术单元提案。最初的TinyStories前缀结果使用了公开的1M字节前缀、无偏乘积投影和仅训练集分词:TriPLU达到平均最佳验证损失1.0637,比紧密匹配的SwiGLU降低了0.0380(约3.4%相对降低)。它还在所有三个随机种子中达到了验证损失目标1.10和1.08;SwiGLU在两个种子中达到1.10,在无一种子中达到1.08。Byte-BPE运行增加了验证选择的检查点、保留集评估、WikiText-2原始文本复现和PMI切片诊断。贡献包括:
- • TriPLU,一种直接乘积前馈网络设计,用学习投影的三线性乘积替换一元前馈网络激活;
- • 一项紧密匹配的TinyStories前缀比较,显示在相同词元预算下TriPLU比SwiGLU达到更低的最佳验证损失;
- • 与分词器无关的Byte-BPE验证集和保留集比较,涵盖TinyStories和WikiText-2原始文本;
- • 相邻词PMI切片,测试增益是否与已见的共现上下文一致;
- • 深度和乘积阶消融实验,展示三阶乘积分支在何处有效,以及二阶或四阶在何处较弱;
- • 学习率和归一化诊断,揭示优化敏感性,但不将其作为主要结论;
- • 算术和门控乘积诊断,将语言建模证据与机制证据分离;
- • 一个可复现的微型GPT基准记录,包含保留集、分词器和FLOP匹配方面的注意事项。

## 2 相关工作
#### Transformer前馈网络变体与门控
GLU风格的Transformer前馈网络是最接近的主流先例。Shazeer (2020) 发现GEGLU和SwiGLU在所研究的序列到序列设置中优于ReLU或GELU前馈网络,这促使SwiGLU作为主要基准。我们的直接乘积前馈网络与SwiGLU的侧重点不同。SwiGLU通过另一个流的非线性变换来门控一个投影流:\[\operatorname{SwiGLU}(x) = \operatorname{SiLU}(W_{g}x) \odot W_{v}x. \quad (1)\] TriPLU,本文研究的主要直接乘积分支,用三个学习投影的乘积替换前馈网络激活路径:\[\operatorname{FFN}_{\mathrm{TriPLU}}(x) = W_{o} \alpha (W_{u}x \odot W_{v}x \odot W_{w}x). \quad (2)\] 在主要字符级后续实验中,\(\alpha\)是一个在探索性筛选中选定的固定标量分支增益,然后在修正运行中保持固定;第3节给出了针对具体运行系列的缩放设置。这种纯乘积形式测试了三线性交互,没有一元激活;比较是在强大的门控前馈网络和更显式的直接乘积偏向之间进行的。

#### 乘积单元神经网络
经典的乘积单元将输入的幂相乘,而不是加权求和输入。它们表达能力强但难以训练。用于回归和预测的进化乘积单元网络表明,显式乘积可以捕捉非线性交互,包括在《神经计算与应用》的预测研究中。TriPLU则在AdamW训练的Transformer前馈网络中插入了一个逐坐标的三线性乘积分支。最近的激活设计工作表明,前馈网络非线性仍然影响语言模型预训练损失和稳定性,包括激活的混合、PolyGLU、PowLU和SSLU。这些论文将非线性视为活跃的设计变量,但在微型、紧密匹配的训练预算下,未隔离出显式的高阶乘积。

#### 乘法交互
Jayakumar等人 (2020) 将门控、注意力、超网络和动态层视为丰富可表示函数类的乘法交互。我们并未将乘法引入神经网络;我们是在受限的微型语言模型设置中测试一种直接乘积前馈网络。Li等人 (2026) 报告,乘积单元残差网络有助于特征交互回归,同时对优化、初始化和残差稳定性仍然敏感。我们的Transformer分支避免了对数域输入乘积单元,但相同的稳定性问题出现在缩放和种子方差结果中。

#### 神经算术模块
神经算术模块是显式乘法和类幂计算的最接近先例。NALU通过学习门控组合算术操作;NAU/NMU和NPU以更强的约束针对加法/减法、乘法和幂运算。Madsen & Johansen (2022) 综述了这些模块作为系统化算术和逻辑泛化的工具。这些文献影响了我们的消融实验:对数域有符号幂分支和整数幂变体例证了单项式假设,但它们表现不如学习投影的直接乘积。这与算术模块的经验相符,即显式算术偏向带来了优化、符号、零值处理和稳定性挑战。Qiu等人 (2024) 警告,微型Transformer的整数乘法失败涉及进位处理和中间结果缓存,而不仅仅是访问乘法操作。因此,我们将`arithmetic_late`作为次要:较低的算术损失是机制证据,但语言建模损失决定了主要主张。

#### 微型语言模型与TinyStories
TinyStories研究极小的语言模型能否从简化的合成故事中产生连贯的英语。WikiText-2原始文本提供了第二个具有不同文本统计信息的语料库。我们的实验是低预算的前缀或采样检查,而非全周期广泛的预训练。最近的BabyLM工作提供了一个临近的低资源参考点。Haller等人 (2025) 在一个紧凑的Qwen风格基准中使用了SwiGLU前馈网络,这使得强大的门控前馈网络基准在此处尤为重要。

## 3 方法
我们使用一个紧凑的minGPT风格PyTorch模型从头训练小型仅解码器Transformer,该模型具有因果自注意力、残差块、层归一化和可配置的前馈网络。匹配的公开TinyStories重跑使用仅训练集字符分词;仅验证集字符映射到`.`。共享设置对TinyStories前缀运行使用上下文长度64,对算术使用32,2层,2头,嵌入宽度96,dropout 0.0,AdamW优化器,学习率0.0003,权重衰减0.1,批大小32。在每组内,固定形状、分词器、数据划分、优化器、词元预算和种子,并通过调整前馈网络宽度来匹配参数数量。

#### 基准模型
加宽的GELU基准是:\[\operatorname{FFN}_{\textsc{gelu}}(x) = W_{o} \operatorname{GELU}(W_{i}x). \quad (3)\] 在匹配运行中,隐藏宽度被加宽到480。强大的门控基准是:\[\operatorname{FFN}_{\textsc{swiglu}}(x) = W_{o} (\operatorname{SiLU}(W_{g}x) \odot W_{v}x), \quad (4)\] 在匹配运行中隐藏宽度为322。

#### 直接乘积前馈网络
唯一的架构改变是用学习投影的逐元素乘积替换前馈网络隐藏激活。一个直接乘积分支将隐藏状态投影到两个、三个或四个流,并在输出投影前逐元素相乘它们。最简单的分支是二阶的:\(p_{2}(x) = W_{u}x \odot W_{v}x. \quad (5)\) TriPLU使用的三阶直接乘积是:\(p_{3}(x) = W_{u}x \odot W_{v}x \odot W_{w}x, \quad (6)\) 四阶乘积阶控制是:\(p_{4}(x) = W_{u}x \odot W_{v}x \odot W_{w}x \odot W_{z}x. \quad (7)\) TriPLU仅使用缩放的三阶乘积分支:\[\operatorname{FFN}_{\mathrm{TriPLU}}(x) = W_{o} \alpha p_{3}(x). \quad (8)\] 标量\(\alpha\)是一个分支增益超参数,用于使乘积分支在数值上与标准前馈网络激活相当;我们将其视为一个稳定性参数。这个纯乘积前馈网络没有一元非线性,但它是非线性的,因为三个线性投影的乘积对输入是三次的:\(p_{3}(cx) = c^{3} p_{3}(x) \neq c p_{3}(x) \quad (9)\) 对于一般的\(c\)。我们保留\(W_{o}\)将乘积分支宽度映射回残差宽度。乘积投影是无偏的。公开数据的TriPLU设置,记录为`triple_prod`,使用分支宽度242和固定的\(\alpha=5.0\);该增益在探索性筛选中选定,并在修正的公开数据后续运行前固定。`double_prod`使用宽度322,`quad_prod`使用宽度193和固定的\(\alpha=25.0\)。缩放设置仅在后续的Byte-BPE诊断中不同。未归一化的Byte-BPE `TriPLU`使用分支宽度576和一个可学习的每层乘积缩放,初始化为5.0。归一化的`TriPLU`首先对乘积分支应用RMS归一化,然后应用固定的乘积缩放0.1。我们明确报告这些设置,因为归一化扫描测试的是缩放控制,而非在所有运行中使用单一固定\(\alpha\)的架构。

#### 负消融实验
我们还测试了受神经算术模块启发的对数域幂单元:\(m(x) = \exp(\operatorname{clip}(W_{e} \log(\|x\| + \epsilon), -c, c)). \quad (10)\) 这些单元可以表示学习到的单项式,但比直接乘积更难优化。硬符号、整数幂和注意力侧乘积消融也较弱,因此我们将它们视为负诊断。

#### 参数匹配
比较在每个特定数据集的词表内匹配参数数量。在主要的TinyStories 1M字节后续中,SwiGLU和`double_prod`各有282,240个参数,TriPLU有282,624个,`quad_prod`有282,048个。剩余的576参数差异来自离散的分支宽度选择,因此我们称之为接近而非精确匹配。

## 4 实验设置
#### 数据集
`tinystories_lite`是一个本地确定性故事风格语料库,仅用于廉价筛选。公开数据运行使用`roneneldan/TinyStories`的确定性前缀。主要的字符级设置`public_tinystories_1m`使用请求的1,000,000训练字节和200,000验证字节,解码为UTF-8并截断到最后一个完整换行符。`arithmetic_lite`是一个次要的合成字符诊断,包含加法和乘法字符串。Byte-BPE扩展使用了TinyStories和WikiText-2原始文本。

#### 分词器无关的Byte-BPE协议
扩展实验增加了仅训练集的Byte-BPE比较,以便结果不局限于字符分词。它们使用8层仅解码器模型,嵌入宽度256,4头,块大小128,dropout 0.0,请求的词表大小512,AdamW优化器,权重衰减0.1,梯度裁剪1.0,种子1-3。匹配的`SwiGLU`分支使用隐藏大小768;未归一化的`TriPLU`使用乘积分支大小576和一个可学习的乘积缩放,初始化为5.0。我们使用两个Byte-BPE角色:一个初始的100M字节Ti

相似文章