一个反直觉的现场:单卡微调 Gemma-2 2B——按理说权重 bf16 才 5 GB 出头——显存监控却在每个 step 的前向末尾拉出一个十几 GB 的尖峰,然后 OOM。罪魁不是注意力,也不是什么 KV cache(训练时根本没有),而是最不起眼的一步:算 loss 之前的那个 logits 矩阵。
这篇讲清楚三件事:这堵「显存墙」是怎么砌起来的;两个专门拆墙的开源方案——LinkedIn 的 Liger Kernel 和 Apple 的 Cut Cross-Entropy(CCE)——各自的刀法;以及怎么在自己的训练里用上它们,包括很多人关心的问题:华为昇腾 NPU 上能不能用。
先正个名。这两样东西经常在训练框架的选项里并排出现(transformers 的 use_liger_kernel、LLaMA-Factory 的 enable_liger_kernel、axolotl 的 CCE 插件……),于是常被连读成「Liger 的 CCE」。其实是两家:Liger Kernel 是 LinkedIn 2024 年开源的一套面向 LLM 训练的 Triton 算子库,其中管 loss 的算子叫 FusedLinearCrossEntropy(FLCE);CCE 则出自 Apple 的论文《Cut Your Losses in Large-Vocabulary Language Models》(ICLR 2025 口头报告),配套开源库 ml-cross-entropy。两者对付的是同一堵墙,思路却分属两个段位:FLCE 是「分小块算,算完就扔」,CCE 是「干脆不让 logits 落地」。两家至今也没合流:Liger 仓库里「把 CCE 方法并进来」的讨论(issue #391)从 2024 年底一直开到今天。
一、这堵墙是怎么砌起来的
Transformer 里绝大多数激活的形状都是 N × H:N 是这一批的 token 总数(batch × 序列长),H 是隐藏维。唯独最后一步,lm_head 把每个位置的表示投到整个词表上:
hidden 态 [N × H] @ lm_headᵀ [H × V] → logits [N × V]
问题出在 V 身上。词表这几年一路通胀:Llama-2 是 3.2 万,Llama-3 涨到 12.8 万,Qwen2.5 是 15 万级,Gemma-2 干脆到了 25.6 万——多语言、代码符号都要塞进去,词表越大、同样文本切出的 token 越少,推理越划算。但训练时的账单也跟着来了:以 Gemma-2 为例,H = 2304、V = 256128,最后这一步把激活宽度放大了 111 倍。
更糟的是这个矩阵不止存一份。主流实现里,logits 由 bf16 矩阵乘算出后,为了在几十万类上做数值稳定的 softmax,要先 .float() 升成 fp32 副本,log_softmax 再生成一份中间结果;反向传播时还得为它配一份同样大小的梯度。CCE 论文在 Gemma-2 2B 上实测(一批 8192 个 token):loss 这一层前向要 24 GB,反向再要 16 GB——论文摘要的原话是,对小模型,这一层吃掉的显存「比模型其余部分加起来还多一个数量级」。
顺带回答一个常见疑问:为什么推理没这个问题?因为自回归解码每步只需要最后一个位置的 logits,矩阵是 1 × V,几百 KB 的事。训练是每个位置都要算 loss,N 一下子变成几万。
二、Liger FLCE:分块、融合、原地覆盖
Liger 的方案建立在两个观察上。
观察一:交叉熵的梯度几乎是免费的。 对 softmax 交叉熵,梯度有个著名的解析形式:
∂loss/∂logits = softmax(logits) − onehot(target)
算 loss 本来就要算 softmax 的分母,所以梯度是顺手就能得到的副产品。而且 loss 是整张计算图的终点,不用等反向传播「传回来」什么——前向算 loss 的同时就可以把这份梯度做掉(严格说还差乘一个上游标量,留到真正的 backward 补上即可)。
观察二:交叉熵按行独立。 第 i 个 token 的 loss 只依赖 logits 的第 i 行。既然行与行互不相干,就没必要让整个 N × V 同时存在——切成小块,一块一块过。
于是 FLCE 的流水线是这样(对应下图):把 N 行切成若干 chunk;每个 chunk 先做小矩阵乘得到一小块 logits;Liger 的 Triton CE kernel 用 online softmax 算出这一块的 loss,并且把 ∂logits 原地写回 logits 块占用的内存(连第二份都不分配);接着立刻把梯度「回投」——∂X 的对应行等于 ∂logits_c @ W,∂W 跨块累加——这块 logits 就功成身退,显存归还,换下一块。
切多大有讲究。源码里的选择是让一块 logits 的体积约等于输入本身:块数取 ⌈V/H⌉(向 2 的幂取整),词表越大切得越碎。源码注释里的例子:N = 16384、V = 32000、H = 4096 时切 8 块、每块 2048 行。这样峰值显存从 N × V 降到 c × V ≈ N × H——跟词表大小基本脱了钩。
代价也要说清楚:矩阵乘被拆小了,kernel 启动次数变多;梯度在前向就算,纯推理场景毫无收益(所以它只该用在训练)。词表不大时,这套操作反而可能比一把梭的大矩阵乘慢。Liger 的定位是算子「全家桶」:除了 FLCE,还有 RMSNorm、RoPE、SwiGLU/GeGLU 等训练算子,官方基准(LLaMA-3 8B,8 × A100,FSDP,bs=8,bf16)给出的端到端数字是吞吐 +20%、显存 −60%,Hugging Face 原版 4K 上下文就 OOM 的配置,加上 Liger 能撑到 16K。同一套「分块 + 梯度前置」的思路还被推广到了后训练:DPO、ORPO、SimPO 等偏好损失的 chunked 版本,README 称最多省 80% 显存。
三、CCE:连一小块都不给你算
Liger 的答案是「小块小块算」,Apple 的追问是:这些 logits 真的需要存在于显存里吗?
把交叉熵拆开看(对数下的 softmax 展开):
loss_i = −log softmax(logits_i)[y_i]
= LSE(logits_i) − logits_i[y_i] 其中 LSE(z) = log Σ_v e^(z_v)
两项各有各的算法:
- 分子项
logits_i[y_i]:正确 token 的那一个 logit。不需要整行——它就是x_i和W里第y_i列的点积。整个 batch 就是 N 次「按标签索引」的点积(indexed matmul),输出 N 个数,显存 O(N)。 - 分母项 LSE:确实要遍历全词表。但「遍历」不等于「保存」——Flash Attention 早就示范过:把大矩阵切成瓦片搬进 SRAM,片上算完就地归并,只把归并结果写回显存。
CCE 的 linear-LSE kernel 正是这么干的:把 (N, V) 平面切成瓦片,每个 GPU block 把对应的 x 块、w 块搬进 SRAM,片上做小矩阵乘得到 logits 瓦片,当场做块内 LSE,再通过带自旋锁的原子操作并入全局 LSE 向量。在线归并用的还是那套经典的 running max 技巧:
m′ = max(m, m_blk)
s′ = s · e^(m − m′) + s_blk · e^(m_blk − m′)
LSE = m′ + log s′
最终写回 HBM 的只有长度 N 的 LSE 向量:8192 个 token 也就 32 KB。所以论文表格里,loss 计算的显存那一格写的是——1 MB。
反向同样不落盘:∂logits = softmax − onehot,而 softmax = e^(logits − LSE),瓦片里重算一遍 logits、乘一乘就有了。这里还藏着 CCE 最漂亮的一招——梯度过滤。训练用 bf16 时,凡是 softmax 值小于 2⁻¹² 的项,累加进梯度也会被舍入吞掉,算了等于白算;而实测中一行 softmax 里「舍入后非零」的元素不到 0.02%(大约排到第 50 名之后的 token,概率就掉穿了阈值)。于是整块都低于阈值的瓦片直接跳过,配合词表排序(前向顺手统计每个词的平均 logit,反向把「热」词排进相邻的块,跳过得更整齐),反向快了 3.5 倍。
论文在 Gemma-2 2B(8192 token,A100 80GB)上的对比表值得放在一起看:
| 实现 | loss 计算显存 | 反向额外显存 | 前向 + 反向耗时 |
|---|---|---|---|
| PyTorch 朴素实现 | 24 GB | 16 GB | 208 ms |
| torch.compile | 4 GB | 12 GB | 143 ms |
| torchtune 分 8 块 | 8 GB | 1.6 GB | 169 ms |
| Liger FLCE(论文当时版本) | 1.5 GB | — | 304 ms |
| CCE | 1 MB | 1.2 GB | 145 ms |
两点解读。其一,CCE 反向剩下的那 1.2 GB,大头是分类头权重梯度 ∂W 本身(2304 × 256128 的 bf16 矩阵就是约 1.1 GB)——这是训练躲不掉的部分,说明 loss 层的「额外」开销已经贴到理论地板。其二,表里 Liger 的耗时是 CCE 论文 2024 年底测的旧版本,此后 Liger 一直在迭代,具体数字别抠,看路线差异就好:分块把显存降一个量级,片上计算再降三个量级。
精度呢?CCE 是确定性算法,论文给了训练曲线:与朴素实现的收敛完全重合。对数值特别敏感的场景有 cce_kahan 变体(Kahan 求和);预训练场景官方建议用关闭词表侧梯度过滤的 cce_kahan_full_c——冷门 token 在预训练里也要更新,不宜过滤。
四、实操
Liger Kernel
pip install liger-kernel
最省事的是 transformers Trainer 一个开关(TRL 的 SFTConfig 继承同样的参数):
from transformers import TrainingArguments
args = TrainingArguments(
...,
use_liger_kernel=True,
# 可选:细粒度控制交给 liger_kernel_config
# liger_kernel_config={"rope": True, "swiglu": True,
# "fused_linear_cross_entropy": True},
)
或者在自己的代码里 patch 模型(也可以用 AutoLigerKernelForCausalLM 一步到位):
from liger_kernel.transformers import apply_liger_kernel_to_qwen2
apply_liger_kernel_to_qwen2(
rope=True, rms_norm=True, swiglu=True,
fused_linear_cross_entropy=True,
)
model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen2.5-7B")
写自定义训练循环的话,直接用低阶 API(注意签名是权重在前):
from liger_kernel.transformers import LigerFusedLinearCrossEntropyLoss
loss_fn = LigerFusedLinearCrossEntropyLoss()
loss = loss_fn(lm_head.weight, hidden_states, targets) # 全程不出现 logits
loss.backward()
框架开关:LLaMA-Factory 是 enable_liger_kernel: true,axolotl、SWIFT(ms-swift)等也有对应集成。
Cut Cross-Entropy
pip install "cut-cross-entropy @ git+https://github.com/apple/ml-cross-entropy.git"
要求 PyTorch 2.4+、Triton 3.0+、Ampere 或更新的 GPU。核心就一个函数:
from cut_cross_entropy import linear_cross_entropy
# shift=1:自动做因果 LM 的错位对齐
loss = linear_cross_entropy(hidden_states, lm_head_weight, labels, shift=1)
现成的 transformers 模型可以直接 patch(README 列出支持 Llama、Phi3、Mistral、Gemma2 等家族):
from cut_cross_entropy.transformers import cce_patch
model = cce_patch(model)
impl 参数可选实现:cce(默认,最省显存)、torch_compile(README 的说法是通常最快但最费显存,也是非 NVIDIA 环境的后备)、cce_kahan / cce_kahan_full_c(更高数值精度 / 预训练用)。另外它原生支持词表并行(VocabParallelOptions),大规模预训练也能接。框架侧:axolotl 有独立的 CutCrossEntropyPlugin(cut_cross_entropy: true),unsloth 则内置了自家维护的 CCE fork。
手写版:分块交叉熵,顺便讲透 use_reentrant
不想引入任何依赖,或者目标硬件连 Triton 都没有?「分块」这个思想本身用纯 PyTorch 就能表达,十几行:
import torch
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
def chunk_loss(h_c, weight, y_c):
logits = (h_c @ weight.T).float() # 只把这一小块升成 fp32
return F.cross_entropy(logits, y_c, reduction="sum")
def chunked_ce(hidden, weight, labels, n_chunks=8):
total = hidden.new_zeros((), dtype=torch.float32)
for h_c, y_c in zip(hidden.chunk(n_chunks), labels.chunk(n_chunks)):
total = total + checkpoint(chunk_loss, h_c, weight, y_c,
use_reentrant=False)
return total / (labels != -100).sum()
(示意骨架:假设 labels 已做过因果错位;ignore_index 用的是 F.cross_entropy 的缺省 −100。)
这就是 torchtune 分块方案的骨架(上一节对比表里 8 GB / 1.6 GB 那一行)。它省显存的关键不在 chunk,而在 checkpoint。autograd 默认会把反向要用的中间量——这里是每一块的 log_softmax 结果——全部保存到反向阶段;只切块不重算,前向峰值确实错开了,但 8 份小 logits 攒到反向时还是一整个 N × V,等于白切。包上 torch.utils.checkpoint 语义才完整:这段前向不保存中间量,只留下每块一个标量 loss;反向传到哪块,就重算哪块的 logits——「任何时刻只有一块活着」在反向也成立了。代价是 logits 的矩阵乘要算两遍(前向一遍、反向重算一遍),用一份额外计算换一个数量级的显存,通常很划算。
于是绕不开那个用过梯度检查点的人都见过的参数:use_reentrant。它选的不是「开或关」,而是 checkpoint 的两台引擎,差别大到官方要求必须显式二选一——PyTorch 文档明说从 2.9 起不传这个参数直接抛异常,推荐值是 False。
use_reentrant=True 是旧引擎(可重入)。 前向整段在 torch.no_grad() 下执行,完全不记录计算图,只把输入存下来;反向传到这一段时,重新执行一遍前向,并在反向过程中再次进入 autograd 引擎跑一个嵌套的 backward——「可重入(reentrant)」说的就是这次重入。嵌套引擎带来一串文档明载的限制:整段必须从头到尾完整重算,一步不能省;只支持 .backward() 一种反向方式,torch.autograd.grad() 不能用;要求输入和输出中至少各有一个 requires_grad=True 的张量,不满足时梯度会静默断掉;藏在 list、dict 等嵌套结构里的张量不被视为参与求导;段内还不能出现 detached 张量。
use_reentrant=False 是新引擎(非重入)。 不再嵌套,前向正常记图,但把「本应保存的中间激活」换成占位符;反向第一次真正用到某个激活时才触发重算,而且算到需要的那个就停(签名里如今有独立的 early_stop 参数,默认开启)。旧引擎的限制全部解除,还附赠两件调试工具:determinism_check 默认核对重算出的张量与当初保存的形状、dtype、设备是否一致;debug=True 能吐出两次前向的完整算子轨迹供比对。平时在 Trainer 里写的 gradient_checkpointing_kwargs={"use_reentrant": False},说的就是同一个参数。新代码没有理由选 True。
最后把三条线收拢:手写版需要 checkpoint,是因为它活在 autograd 的语义里——「用重算换显存」只能靠这个 API 表达。Liger FLCE 没有这个参数,因为它的梯度在前向就算完了,压根不存在「反向重算」;CCE 的重算发生在 kernel 内部的 SRAM 里,同样不经过 autograd。三者是同一个思想在三个抽象层级上的落地。
怎么选
- 想省事:框架开关打开 Liger,RMSNorm、RoPE、SwiGLU 连带 loss 一起融了,端到端收益稳。
- 词表巨大、序列极长、显存贴地飞行:CCE,loss 层显存直接归零。
- 不能用 Triton 的环境,或想先搞懂再上轮子:手写分块 + checkpoint,一个量级的节省唾手可得。
- 组合拳:两者 patch 的部位不同,可以把 Liger 的
fused_linear_cross_entropy=False关掉、loss 交给 CCE,其余算子照融。 - 提醒三句:这些优化对推理没有意义(推理只算最后一个位置);小词表模型收益有限;它们与梯度检查点、FSDP、Flash Attention 正交,可叠加。
五、华为昇腾 NPU 能用吗?
Liger Kernel:能,而且是上游官方支持,不是 fork 补丁。 时间线很清晰:2025 年 11 月华为侧在 Liger 仓库发起 RFC(issue #954),一周后首个 NPU 支持 PR 合并;v0.8.0(2026 年 4 月)的 release note 正式官宣支持昇腾后端;最新的 v0.8.1(2026 年 7 月)把依赖定在 torch == 2.7.1 + torch_npu == 2.7.1 + triton-ascend == 3.2.1。这不是纸面适配:仓库里有独立的昇腾后端目录 ops/backends/_ascend/,二十多个 NPU 专用 kernel——交叉熵和 FLCE 都在其中,还按昇腾的片上资源调小了分块上限(主干代码里那句 MAX_FUSED_SIZE = 2048 if infer_device() == "npu")——NPU 环境下 liger_kernel.ops 会自动分发到这套实现。安装时唯一的特殊动作是 triton-ascend 要从昇腾镜像装:
pip install triton-ascend==3.2.1 --extra-index-url https://triton-ascend.osinfra.cn/pypi/simple
官方 roadmap(issue #969)的口径:22 个 kernel 在 Atlas 900 A2 POD(即 910B/A2 系列)和 Atlas 800I A2 上通过了全量 transformers 集成测试,A3 的支持也已合并。保留意见也有两条:NPU CI 因政策原因跑在贡献者 fork 上;2026 年 8 月还有 open PR 在修 NPU 上 FLCE 反向的 fp32 梯度问题。结论是「能用、在快速打磨中」,版本组合照官方 pin 的来,别自己混搭。
框架层的打通情况(截至 2026-08):
| 框架 | 昇腾 + Liger 状态 |
|---|---|
| transformers / TRL | use_liger_kernel=True 只检查包、不检查设备,NPU 直接走通 |
| LLaMA-Factory | enable_liger_kernel: true 开箱即用(issue #10386,Atlas 900 A2 实测) |
| verl | 官方合并(PR #6244):卸载 triton、改装 triton-ascend,model.use_liger=true |
| ms-swift | 参数在(--use_liger_kernel),NPU 官方验证未落地 |
这一切的底座是 triton-ascend——把 Triton 编译栈接到 CANN 上的项目,如今已从华为自留地搬进 Triton 官方组织(triton-lang/triton-ascend),3.2.2 于 2026 年 7 月底发布,自述覆盖约 85% 的 Triton Python API。Liger README 那句「我们的 kernel 继承 Triton 提供的全部硬件兼容性」,在昇腾身上兑现了大半。生态倒也没到丝滑:ms-swift 社区就有 PR 宁可用纯 PyTorch 分块实现,理由写得直白——绕开 triton-ascend 的编译器问题。
CCE:官方不支持。 仓库要求 Ampere 或更新的 NVIDIA GPU,代码里唯一的非 Triton 分支是给 macOS 准备的 torch_compile 路径,issue 和 PR 里检索 NPU/Ascend 是零命中。torch_compile 变体在 torch_npu 的编译支持下理论可行,但中英文都搜不到跑通的公开记录——别当成能用。
昇腾上想省 loss 显存,现实梯队是:
- Liger FLCE:官方路线,见上;
- CANN 自带融合算子:
torch_npu.npu_cross_entropy_loss(op-plugin 7.0 起,Atlas A2/A3),融合的是 CE 本体(含 z-loss),但输入就是 logits——矩阵还得先物化,省的是 fp32 副本和中间量那几份,且单次上限 20 万行; - 手写分块版:上一节的纯 PyTorch 实现不依赖 Triton,torch_npu 后端直接跑,ms-swift 社区那个 PR 就是同款思路。
值得盯一眼的苗头:op-plugin 的新分支里已经出现了 linear + CE 一体的融合算子(npu_fused_cross_entropy_loss_with_max_sum 一族,反向签名带上了 lm_head 权重),只是还没查到公开的上层封装。
尾声
把三代方案排成一行,其实是同一个故事讲了三遍:
少存几份(治理 fp32 副本)→ 分块存(Liger FLCE)→ 干脆不存(CCE)
这和 Flash Attention 的故事一模一样,只是舞台从注意力矩阵换成了 logits 矩阵:HBM 才是稀缺资源,能在片上算完的,就别落盘。 词表还会继续变大,但训练显存的账单,从此可以不跟着涨。
参考
- Liger Kernel:github.com/linkedin/Liger-Kernel,论文 arXiv:2410.10989
- Cut Cross-Entropy:github.com/apple/ml-cross-entropy,论文 arXiv:2411.09009(ICLR 2025)
- triton-ascend:github.com/triton-lang/triton-ascend(分发走昇腾 PyPI 镜像 triton-ascend.osinfra.cn)
- PyTorch 梯度检查点文档(use_reentrant 两种实现的差异):docs.pytorch.org/docs/stable/checkpoint.html