重新思考标记预测:树结构扩散语言模型
Rethinking Token Prediction: Tree-Structured Diffusion Language Model
AI总结:
本文提出树结构扩散语言模型,通过利用标记间的内在结构,减少参数和内存使用,实现在有限资源下提升性能。
AI中文摘要:
离散扩散语言模型已逐渐成为自回归语言模型的有力替代品,但在有限的参数和内存预算下高效训练仍具挑战性。现代架构大多基于全词汇标记预测层,这占用了模型参数的大量比例(例如,在小型DiT风格设计中超过20%),并且常常主导峰值GPU内存使用。这导致在受限训练资源下参数和内存的使用效率低下。为了解决这一问题,我们重新审视显式全词汇预测的必要性,并利用标记之间的内在结构,构建树结构扩散语言模型。具体而言,我们通过预构建的词汇树中标记的祖先节点对应的中间潜在状态来建模扩散过程。这种树结构分解指数级降低了分类维度性,使预测头的尺寸变得可以忽略不计,并使参数能够重新分配以加深注意力块。实验表明,在相同参数预算下,我们的方法将峰值GPU内存使用量减少了一半,同时达到了最先进的离散扩散语言模型的困惑度性能。
英文摘要:
Discrete diffusion language models have emerged as a competitive alternative to auto-regressive language models, but training them efficiently under limited parameter and memory budgets remains challenging. Modern architectures are predominantly based on a full-vocabulary token prediction layer, which accounts for a substantial fraction of model parameters (e.g., more than 20% in small scale DiT-style designs) and often dominates peak GPU memory usage. This leads to inefficient use of both parameters and memory under constrained training resources. To address this issue, we revisit the necessity of explicit full-vocabulary prediction, and instead exploit the inherent structure among tokens to build a tree-structured diffusion language model. Specifically, we model the diffusion process with intermediate latent states corresponding to a token's ancestor nodes in a pre-constructed vocabulary tree. This tree-structured factorization exponentially reduces the classification dimensionality, makes the prediction head negligible in size, and enables reallocation of parameters to deepen the attention blocks. Empirically, under the same parameter budget, our method reduces peak GPU memory usage by half while matching the perplexity performance of state-of-the-art discrete diffusion language models.