跳到主要内容

数据截至 (上游 commit 1fe27b1b53f3)

03 · 融合交叉熵:让 logits 从不存在

这章回答:为什么 Unsloth 训长序列不 OOM。答案:训练时完整的 logits 张量从头到尾就没 被创建过

1. 小问题

Causal LM 训练的最后一步是把 hidden states 乘上 lm_head 得到 logits,形状 (batch, seq_len, vocab_size)。这笔账很吓人:

  • 8B 模型词表 128K、batch×seq = 8192:logits = 8192 × 128256 × 2 字节 ≈ 2 GB;
  • 再加上反向要存 logits(算 softmax 梯度),峰值轻松翻倍;
  • Gemma 这类 256K 词表更夸张。

而 logits 只是中间量——我们真正要的只有 loss 一个标量,和 logits 对 loss 的梯度。

2. 思路 / 直觉

交叉熵有个干净的闭合式:

  • 前向:CE = logsumexp(logits) - logit[label];
  • 反向:dCE/dx = softmax(x),在 label 处再减 1。

两者都不需要同时看到全部 logits,只要按块流式处理:算 logsumexp 可以先逐块算局部 logsumexp 再合并;softmax 用 exp(x - logsumexp) 逐块还原。这就是「融合交叉熵」的数学 地基——kernel 注释里把推导完整写了一遍(unsloth/kernels/cross_entropy_loss.py:51-67:214-230)。

3. 两层结构:kernel 在 kernels/,分块调度在 unsloth_zoo

labels ──┐

CausalLM_fast_forward ──有 labels?──► unsloth_fused_ce_loss(unsloth_zoo)
(unsloth/models/llama.py:1386) │ 按显存目标选 chunk 大小

逐块:hidden @ lm_head 的列块 → logits 小块

Fast_CrossEntropyLoss(本章主角,unsloth/kernels/)

每行一个 loss

读法: 上层决定「切多大块」(调度,在依赖包 unsloth_zoo);下层决定「块内怎么算」 (kernel,在本仓库)。两层加起来,任意时刻显存里只有一块 logits。

4. kernel 层:Fast_CrossEntropyLoss

前向

Fast_CrossEntropyLoss.forward(unsloth/kernels/cross_entropy_loss.py:288-377)按词表 大小分两路:

词表路径做法
≤ 65536(MAX_FUSED_SIZE,unsloth/kernels/utils.py:20)_cross_entropy_forward(:35)每行一个 program,整行读入,一次算 logsumexp 和 loss
> 65536(如 Gemma 256K)_chunked_cross_entropy_forward(:114)按 65536 列分块,每块算局部 logsumexp,最后 torch.logsumexp(dim=1) 合并(:368-370)

单行版的核心三步(:93-107):先对 logits 做可选的 scaling(Gemma2 softcap / Cohere logit_scale,靠 triton.heuristics 开关编进 kernel,:107-113),再 c = max(x) 稳化算 logsumexp,最后 loss = logsumexp - x_label;label 为 -100(padding)直接记 0。

反向:梯度写回 logits 缓冲区

_cross_entropy_backward(:202-285)按 (n_rows, n_blocks) 网格启动,每块:

# 示意,非源码 —— 反向核心
y = exp(x - logsumexp) # softmax,用前向存的 logsumexp
y = where(col == label, y - 1, y)
tl.store(logits_ptr + col, dloss * y) # 写回 logits 原内存

注意最后一行:梯度直接覆盖在 logits 块上(:284),不新开显存。softcap/scaling 的 导数(1 - tanh²)也在块内乘好(:272-276)。

标量汇总

fast_cross_entropy_loss(:421-452)做最后的 / n_items 归一——n_itemslabels != -100 的计数,正好对上 transformers 的 num_items_in_batch(梯度累积下的正确 分母)。

5. 接线:训练时 logits 变成了「一调用就报错」的占位符

接回顶层:CausalLM_fast_forward(unsloth/models/llama.py:1386)在「有 labels 且不要求 返回 logits」时,直接调 unsloth_fused_ce_loss 算 loss 返回(llama.py:1495-1513),然后 把一个 EMPTY_LOGITS 塞进返回值的 logits 槽位(llama.py:1527-1534)。

EMPTY_LOGITS(unsloth/models/_utils.py:3734-3747)是个所有张量方法都被换成 「抛异常」的对象:下游代码一旦真去用 logits,立刻炸出明确错误,而不是悄悄拿到错的形状。

两个由此而来的边界开关:

  • UNSLOTH_RETURN_LOGITS=1 强制走「真的算出 logits」的老路(llama.py:1479);
  • bsz*q_len <= 1024 时本来就更省,Unsloth 仍走融合路径但不强制(llama.py:1481-1484)。

6. 还有一个补丁:patch_loss_functions

除了模型前向这一入口,Unsloth 还直接改 transformers 的 loss 注册表: patch_loss_functions(unsloth/kernels/cross_entropy_loss.py:459-473)把 LOSS_MAPPING 里所有指向原版 ForCausalLMLoss 的别名统一切到 Unsloth 版——就算模型走 的是 HF 标准 loss 路径,落到手的也是融合 kernel。

7. 坑

  • T4 放不下大 kernel:源码里保留着被禁用的 fused_linear_cross_entropy 调用和注释—— T4 共享内存 64KB,块太大直接 OutOfResources(llama.py:1507-1512)。现在的 unsloth_fused_ce_loss 会按显存目标自动选块大小。
  • packed 序列的边界:fused kernel 内部自己做 label 右移,所以 packing 场景要先用 mask_packed_boundary_labels 把跨样本边界的 label 抹成 -100(llama.py:1497-1503), 否则模型会学「预测下一个样本的第一个 token」。
  • 反向依赖前向存的 logsumexp:每行一个 fp32 标量,这是唯一为反向保留的激活—— 也是整套方案「存 O(n) 换 O(n×V)」的精髓。