大词表语言模型中的损失裁剪
机器学习
2025-03-12 v2 计算与语言
摘要
随着语言模型规模不断增大,其词表也随之扩大。这使得LLM训练时的内存占用不成比例地集中在单一层上:损失计算中的交叉熵。交叉熵会构建一个logit矩阵,其条目对应每一对输入token和词表项,对于小模型,其消耗的内存比LLM其余部分总和还要多一个数量级。我们提出了Cut Cross-Entropy (CCE),一种无需将所有token的logit实例化到全局内存即可计算交叉熵损失的方法。相反,CCE仅计算正确token的logit,并即时对所有logit求log-sum-exp。我们实现了一个自定义内核,在闪存中执行矩阵乘法和对词表的log-sum-exp归约,使交叉熵计算的全局内存消耗可忽略不计。效果显著。以Gemma 2 (2B)模型为例,CCE将损失计算的内存占用从24 GB降至1 MB,将分类头的总训练时内存消耗从28 GB降至1 GB。为提高CCE的吞吐量,我们利用softmax的固有稀疏性,提出跳过对梯度贡献可忽略(即低于数值精度)的梯度计算元素。实验表明,在不牺牲训练速度或收敛性的前提下,实现了内存消耗的大幅降低。
引用
@article{arxiv.2411.09009,
title = {Cut Your Losses in Large-Vocabulary Language Models},
author = {Erik Wijmans and Brody Huval and Alexander Hertzberg and Vladlen Koltun and Philipp Krähenbühl},
journal= {arXiv preprint arXiv:2411.09009},
year = {2025}
}
备注
To appear in ICLR 2025 (Oral). Code is available at https://github.com/apple/ml-cross-entropy