一个反直觉的现场:单卡微调 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 倍。

前向的最后一步:形状突然放大hidden 态N × H×lm_head 权重 WᵀH × V=logitsN × Vbf16 一份、升 fp32 一份、log_softmax 再一份…Gemma-2 2B:H = 2304,V = 256128 —— 最后一步把激活宽度放大 111 倍bs 8 × seq 4096(N = 32768):logits 仅 bf16 一份就有 16.8 GB,而模型全部权重 bf16 才约 5 GBCCE 论文实测(N = 8192):loss 前向 24 GB、反向再 16 GB —— 比模型其余部分加起来大一个数量级
前向最后一步的形状账:一步放大 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 就功成身退,显存归还,换下一块。

FLCE 流水线:一次只让一小块 logits 存在hidden 态 XX_cN×H,切成若干块X_c @ Wᵀlogits 块:c 行 × V 列只在此刻存在Liger CE kernel(Triton)online softmax:loss 与 ∂logits 同时出∂logits 原地覆盖Σ → 总 loss∂X逐块写回∂X_c 写回∂W += ∂logits_cᵀ · X_c跨块累加(H×V)释放 logits 块,换下一块峰值显存:一块 logits(c×V ≈ N×H)+ 训练本来就要有的 ∂X、∂W —— 与词表大小基本脱钩块数 ≈ ⌈V/H⌉:词表越大切得越碎。源码注释原例:N=16384、V=32000、H=4096 → 切 8 块、每块 2048 行
FLCE:任何时刻只有一小块 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。

CCE:把 loss 拆成两项,哪一项都不需要完整 logitslossᵢ = LSE(xᵢWᵀ) − xᵢ · w[yᵢ]分子:只挑正确 token 那一列 → N 次点积分母:要看全词表 → 分块在片上算分子:indexed matmulX:N×H·W:V 列里只取第 yᵢ 列=一个数整个 batch:N 次点积、输出 N 个数 —— 显存 O(N)分母 LSE:瓦片进 SRAM,在线归并N×V,从不落盘SRAM 片上x 块 × w 块块内 LSE在线归并(m, s) 原子更新LSE 向量:长度 N ≈ 32 KB写回 HBM 的只有 LSE —— 论文表里 loss 显存那格:1 MB反向:softmax 的稀疏性 → 梯度过滤S = e^(z − LSE),一行 softmax实测非零元 < 0.02%S < 2⁻¹² 的项在 bf16 梯度里加了也会被舍入吞掉 → 整块低于阈值就整块跳过,反向快 3.5×再把词表按平均 logit 排序,让「热」token 聚到相邻块 —— 跳过得更整齐
CCE 前向:分子走索引点积,分母在 SRAM 里分块归并;反向靠稀疏性整块跳过

反向同样不落盘:∂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 GB16 GB208 ms
torch.compile4 GB12 GB143 ms
torchtune 分 8 块8 GB1.6 GB169 ms
Liger FLCE(论文当时版本)1.5 GB—304 ms
CCE1 MB1.2 GB145 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 / TRLuse_liger_kernel=True 只检查包、不检查设备,NPU 直接走通
LLaMA-Factoryenable_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 显存,现实梯队是:

  1. Liger FLCE:官方路线,见上;
  2. CANN 自带融合算子:torch_npu.npu_cross_entropy_loss(op-plugin 7.0 起,Atlas A2/A3),融合的是 CE 本体(含 z-loss),但输入就是 logits——矩阵还得先物化,省的是 fp32 副本和中间量那几份,且单次上限 20 万行;
  3. 手写分块版:上一节的纯 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 才是稀缺资源,能在片上算完的,就别落盘。 词表还会继续变大,但训练显存的账单,从此可以不跟着涨。


参考