数据截至 (上游 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_items 是
labels != -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)」的精髓。