通过联合预测多个未来词元,提升模型表征能力且开销极小。
Efficient Joint Prediction of Multiple Future Tokens
- 设计瓶颈结构,用教师强制机制联合预测多词元。
- 在星图导航任务中显著优于现有方法,实现短时信念状态表征。
- 轻量级改进,适合追求高效表征学习的研究者。
本文提出联合多词元预测(JTP),一种对标准下一个词元预测的轻量级改进,通过联合预测多个未来词元来丰富隐藏状态表示。与以往方法不同,JTP通过精心设计的表示瓶颈,利用教师强制策略引入未来词元信息,使模型在训练中以极小计算开销编码丰富的预测信息。我们证明,JTP能实现短时信念状态表示,而主流多词元预测方法无法做到这一点。在Bachmann和Nagarajan [2024] 提出的合成星图导航任务上,该方法展现出显著性能提升。本研究呈现有前景的初步结果,旨在激发进一步探索。
原文摘要 · Abstract (English)
In this short report, we introduce joint multi-token prediction (JTP), a lightweight modification of standard next-token prediction designed to enrich hidden state representations by jointly predicting multiple future tokens. Unlike previous multi-token prediction approaches, JTP strategically employs teacher forcing of future-tokens through a carefully designed representation bottleneck, allowing the model to encode rich predictive information with minimal computational overhead during training. We show that the JTP approach achieves a short-horizon belief state representation, while popular alternatives for multi-token prediction fail to do so. We demonstrate the effectiveness of our method on the synthetic star graph navigation task from from Bachmann and Nagarajan [2024], highlighting a significant performance improvement over existing methods. This manuscript presents promising preliminary results intended to stimulate further research.
Thank you to arXiv for use of its open access interoperability. PaperDance 不是 arXiv 官方产品;中文卡片由大模型生成,请以原文为准。