[论文解读] RecurrentGemma: Moving Past Transformers for Efficient Open Language Models
RecurrentGemma-2B 引入了一种基于 Griffin 架构的开源语言模型,该架构用线性递推和局部注意力替代了全局注意力,实现了固定大小的状态表示。这使得在长序列上进行推理时速度显著提升且内存占用更低,尽管训练时使用的 token 数量减少了 33%,其性能仍与 Gemma-2B 保持一致。
We introduce RecurrentGemma, a family of open language models which uses Google's novel Griffin architecture. Griffin combines linear recurrences with local attention to achieve excellent performance on language. It has a fixed-sized state, which reduces memory use and enables efficient inference on long sequences. We provide two sizes of models, containing 2B and 9B parameters, and provide pre-trained and instruction tuned variants for both. Our models achieve comparable performance to similarly-sized Gemma baselines despite being trained on fewer tokens.
研究动机与目标
- 开发一种高效、开源的语言模型,使其在长序列上的推理速度和内存效率方面优于 Transformer 模型。
- 证明基于线性递推的模型能够实现与 SOTA Transformer 模型(如 Gemma-2B)相当的性能。
- 通过用固定大小的状态替代 Transformer 的线性增长 KV 缓存,实现长上下文生成。
- 发布预训练版本和指令微调版本的模型,供资源受限环境中的研究与部署使用。
- 通过严格的基准测试和人工评估验证模型的安全性和对齐性,遵循 Gemma 的负责任 AI 实践。
提出的方法
- 采用 Griffin 架构,结合线性递推(RG-LRU)与局部注意力(窗口大小 2048),在不使用全局注意力的情况下建模序列。
- 使用固定大小的状态向量压缩输入序列,消除自回归生成过程中对不断增长的 KV 缓存的需求。
- 对输入嵌入应用可学习的缩放因子(sqrt(模型宽度)),与 Gemma 的设计保持一致,以稳定训练过程。
- 使用专用的 Pallas 内核在 TPU 上实现高效推理,同时提供参考的 PyTorch 实现。
- 在与 Gemma-2B 相同的数据集上训练了 2T 个 token,采用两阶段预训练流程:先是一般性混合数据,再是高质量数据。
- 通过指令微调和一种新颖的 RLHF 算法对模型进行微调,以提升指令遵循能力,使用包含控制标记的定义对话格式。
实验结果
研究问题
- RQ1基于线性递推和局部注意力的循环架构是否能在标准 NLP 基准测试中实现与 Gemma-2B 等 Transformer 模型相当的性能?
- RQ2用固定大小状态替代 KV 缓存,是否能显著提升长序列上的推理速度并降低内存占用,相比标准 Transformer 模型?
- RQ3在仅使用更少 token(2T 对比 3T)的情况下,模型在保持效率的同时,其性能能与更大规模预训练方案相比达到何种程度的匹配?
- RQ4与已建立的指令微调模型(如 Mistral 7B)相比,该模型在人工评估中的表现如何?
- RQ5Griffin 架构能否支持任意长度的生成而无内存限制?其吞吐量如何扩展?
主要发现
- RecurrentGemma-2B 在学术基准测试套件中平均得分为 44.6%,与 Gemma-2B 的 45.0% 相当,表明尽管训练 token 数量减少了 33%,其性能仍具竞争力。
- 在人工评估中,RecurrentGemma-2B-IT 在 1,000 个指令遵循提示上的胜率达到了 43.7%,相较于 Mistral 7B v0.2 Instruct,显示出强大的对齐性和可用性。
- 在 TPUv5e 设备上,RecurrentGemma 在所有序列长度下的推理吞吐量均持续高于 Gemma,且随着序列长度增加无性能下降。
- 在自回归采样过程中,RecurrentGemma 保持了高吞吐量(6k tokens/sec),而 Gemma 因 KV 缓存持续增长导致吞吐量显著下降,尤其在长序列上更为明显。
- 两种模型的提示处理速度相近(约 40k tokens/sec),确认性能优势仅体现在自回归生成阶段。
- 该模型的固定大小状态支持任意长度生成,仅受计算能力和上下文窗口限制,而 Transformer 模型则受限于内存密集型 KV 缓存的增长。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。