短文拼进同一窗口会互相偷看吗?Sequence Packing 的边界掩码
从固定上下文窗口的 padding 浪费出发,手算 sequence packing,分别构造块对角因果注意力、跨文档 loss 屏蔽与位置编号,并给出可验证的 PyTorch 实现。
上一篇确定了不同数据来源应占多少 token 预算。可真正送入模型时,新闻可能只有 180 token,代码文件可能有 3,000 token,而训练内核希望每行长度固定为 。每篇文档单独补齐会把大量算力花在 padding;直接首尾相接又可能让后一篇“看见”毫不相关的前文。
序列装箱(Sequence Packing)要同时解决三个边界:哪些 token 放进同一行、谁能注意谁、哪些相邻对产生 next-token loss。三者不是一张 mask,也不能只插一个 EOS 就假定问题消失。
01 Padding 浪费的究竟是什么?#
设一个 batch 有 行,每行物理长度 ,有效 token 数为 。token 利用率为:
四篇含 EOS 的文档长度为 [4, 2, 3, 3],若每篇都补到 ,有效率只有 。把 [4,2] 与 [3,3] 分别装进两行,则 。
但 Transformer 注意力的主要工作量近似与 成正比。少一半物理行,不只是少存 PAD,也减少了 QK 点积、激活与通信。
02 一行里必须保存哪些元数据?#
flowchart LR
A[带来源和 doc_id 的文档] --> B[tokenize + EOS]
B --> C[packing 算法]
C --> D[input_ids N×L]
C --> E[document_ids N×L]
E --> F[块对角 causal attention mask]
E --> G[跨文档 loss mask]
E --> H[position_ids N×L]
D --> I[Transformer]
F --> I
H --> I
I --> J[logits N×L×V]
G --> K[next-token loss]
J --> Kmermaidinput_ids 只说明 token 是谁,不说明它属于哪篇文档。最小可审计表示应额外保留 document_ids[N,L];PAD 用 -1,真实文档用稳定 id。来源、原文偏移和质量标记可作为旁路元数据,不要塞进模型词表。
03 一个 6-token 包怎样手算?#
两篇文档(均已含 EOS)为:
doc 7: [A, B, EOS] doc 9: [C, D, EOS]
input: [A, B, EOS, C, D, EOS]
docid: [7, 7, 7, 9, 9, 9]
pos: [0, 1, 2, 0, 1, 2]text普通因果 mask 会让 C 看见 [A,B,EOS]。块对角因果 mask 只允许“同文档且 key 位置不晚于 query”:
key → 0 1 2 3 4 5
query 0 ■ · · · · ·
1 ■ ■ · · · ·
2 ■ ■ ■ · · ·
3 · · · ■ · ·
4 · · · ■ ■ ·
5 · · · ■ ■ ■text这是一张由两个下三角块组成的可见图。C 的隐藏状态与 doc 7 无关,即使两篇物理上相邻。
04 Attention Mask 与 Loss Mask 阻止不同泄漏#
对 token 位置 查询 ,允许注意力的条件是:
next-token 标签通常是 。只有上下文 token 与目标 token 属于同一文档时才计损失:
在上例中,B → EOS 应监督,因为 EOS 属于 doc 7;EOS → C 必须忽略。若只做 attention mask 而不做 loss mask,模型仍会被要求从 doc 7 的 EOS 猜 doc 9 的首词。若只做 loss mask,doc 9 的内部预测仍可能借用 doc 7 的隐藏信息。
05 位置编号应重置还是连续?#
两种方案都可能成立,但训练、评测与推理必须一致:
| 方案 | position_ids | 优点 | 风险 |
|---|---|---|---|
| 包内连续 | 0,1,2,3,4,5 | 实现简单 | 后装入的短文总从较大位置开始 |
| 文档内重置 | 0,1,2,0,1,2 | 每篇都像独立样本 | 必须配合文档隔离 mask |
对 RoPE,位置 id 直接决定 Q/K 的旋转角。本文选择文档内重置,使同一篇文档单独运行与装箱运行更容易逐元素对齐。不能只重置位置却保留跨文档注意力:两个文档会出现相同位置坐标并互相可见。
06 装箱算法怎样决定组合?#
给定容量 ,离线数据可用首次适应递减(First-Fit Decreasing):先按长度降序,再把文档放入第一个剩余空间足够的包。它不是最优装箱保证,却比随机相邻稳定。
sort documents by length descending
for document in documents:
for pack in open_packs:
if pack.remaining >= len(document):
append document to pack
break
else:
open a new packtext超长文档不能悄悄丢弃。应明确采用截断、带重叠滑窗或保持文档状态的连续切块,并记录原文区间。在线训练还要限制缓冲区大小,否则“等待更合适的短文”会占满内存并改变采样顺序。
07 用 PyTorch 构造四个训练张量#
import torch
IGNORE = -100
def build_packed_row(documents, doc_ids, block_size, pad_id):
"""documents: list[LongTensor[Li]],每篇已经以 EOS 结尾。"""
if sum(map(len, documents)) > block_size:
raise ValueError("documents exceed block_size")
tokens, owners, positions = [], [], []
for ids, doc_id in zip(documents, doc_ids, strict=True):
tokens.extend(ids.tolist())
owners.extend([doc_id] * len(ids))
positions.extend(range(len(ids)))
pad = block_size - len(tokens)
input_ids = torch.tensor(tokens + [pad_id] * pad) # [L]
document_ids = torch.tensor(owners + [-1] * pad) # [L]
position_ids = torch.tensor(positions + [0] * pad) # [L]
return input_ids, document_ids, position_idspythonBatch 后三者都是 [N,L]。词嵌入输出为 [N,L,D];多头拆分后的 Q/K/V 为 [N,H,L,d],其中 。
08 块对角因果 mask 怎样落到 SDPA?#
当前 PyTorch 2.14 的 scaled_dot_product_attention 中,布尔 attn_mask=True 表示该位置允许参与。自定义块对角因果 mask 已含因果关系,因此调用时设 is_causal=False。
import torch.nn.functional as F
def packed_attention_mask(document_ids):
# document_ids: [N,L]
same_doc = document_ids[:, :, None] == document_ids[:, None, :] # [N,L,L]
valid = document_ids >= 0
causal = torch.ones(
document_ids.size(1), document_ids.size(1), dtype=torch.bool,
device=document_ids.device,
).tril()
allowed = same_doc & valid[:, :, None] & valid[:, None, :] & causal
# PAD query 没有训练意义,但给它开放自身,避免整行均被屏蔽。
eye = torch.eye(document_ids.size(1), dtype=torch.bool,
device=document_ids.device)
allowed |= (~valid)[:, :, None] & eye
return allowed[:, None, :, :] # [N,1,L,L],广播到 H 个头
attn_mask = packed_attention_mask(document_ids)
hidden = F.scaled_dot_product_attention(
q, k, v, attn_mask=attn_mask, dropout_p=0.0, is_causal=False
) # [N,H,L,d]python教学实现显式生成 [N,L,L],便于检查但在长序列上占 内存。生产系统应使用能表达变长/块对角布局的高效 kernel 或元数据接口;不要为了省 padding 又创建一张更昂贵的稠密 mask。
09 Loss Mask 怎样与标签右移对齐?#
def packed_lm_loss(logits, input_ids, document_ids):
# logits [N,L,V];位置 i 预测 input_ids[:, i+1]
pred = logits[:, :-1, :] # [N,L-1,V]
labels = input_ids[:, 1:].clone() # [N,L-1]
same_transition = (
(document_ids[:, :-1] == document_ids[:, 1:])
& (document_ids[:, 1:] >= 0)
)
labels[~same_transition] = IGNORE
return F.cross_entropy(
pred.transpose(1, 2), labels, ignore_index=IGNORE, reduction="sum"
), same_transition.sum()
loss_sum, valid_tokens = packed_lm_loss(logits, input_ids, document_ids)
loss = loss_sum / valid_tokens.clamp_min(1)pythonPyTorch 2.14 的 cross_entropy 会让 ignore_index 目标不贡献梯度;reduction="mean" 也会按未忽略目标平均。这里显式返回和与有效 token 数,是为了多卡或梯度累积时按全局有效 token 归一化,而不是平均各卡的局部均值。
10 最强正确性测试:装箱前后必须等价#
关闭 dropout,把每篇文档单独运行,再与 packed 行对应区间比较:
model.eval()
with torch.no_grad():
packed_logits = model(input_ids, position_ids, attn_mask)
solo_a = model(doc_a[None], torch.arange(len(doc_a))[None], causal_a)
solo_b = model(doc_b[None], torch.arange(len(doc_b))[None], causal_b)
torch.testing.assert_close(packed_logits[0, :len(doc_a)], solo_a[0])
start = len(doc_a)
torch.testing.assert_close(packed_logits[0, start:start+len(doc_b)], solo_b[0])python若不相等,按顺序检查:attention 可见图、position id、padding query、dropout 随机性,再检查是否存在依赖整行统计的自定义层。比较 tolerance 应结合 dtype;先用 float32 建立语义基线。
11 训练流水线还要记录什么?#
每个 pack 至少记录:
- 原始文档 id、来源与 token 区间;
- 装箱算法版本、容量 与文档顺序;
- 有效 loss token 数、padding 数、跨边界屏蔽数;
- 超长文档的截断或切块策略;
- 随机种子、worker/rank 分片和恢复 cursor。
若以“每 step 固定行数”控制训练,packing 提升会让每 step 的有效 token 增加,学习率与总 token 预算因此改变。比较实验必须固定有效 token 或明确报告差异。
12 常见错误与最短调试路径#
| 症状 | 常见原因 | 最短检查 |
|---|---|---|
| packed loss 异常更低 | 后文偷看前文 | 打印一个包的 [L,L] 可见图 |
| 每篇首词 loss 很怪 | EOS→下一篇仍计标签 | 数 doc[i] != doc[i+1] 的标签 |
| RoPE 结果不等价 | position id 未按契约重置 | 打印每篇首尾 position |
| 出现 NaN | PAD query 没有任何可见 key | 检查 mask 每行至少一个 True |
| 吞吐反而下降 | 使用稠密块 mask | profile mask 内存与 attention kernel |
| 恢复后样本变化 | 未保存 packer 缓冲区 | 对比恢复点后前 10 个文档 id |
13 失败场景与相近方法#
Sequence packing 不会减少有效 token 本身的计算,也不会让超长单篇文档突破上下文长度。长度分桶(Length Bucketing)只是让相近长度样本同 batch,仍有 padding;PackedSequence 主要服务 RNN 的变长序列,并不自动给 Transformer 生成块对角注意力。把多篇文档简单 concatenate 后用普通 causal mask,属于连续 token 流训练,不等同于文档隔离 packing。
有些语言模型有意允许跨文档注意力,借 EOS 学习边界。那是另一种训练分布,并非必然错误;但必须明确、做消融,并防止评测样本与训练样本被拼进同一上下文。
14 今天真正需要记住什么?#
- Packing 的目标是减少物理 padding;有效率要按有效 token 与实际 attention 工作量分别测。
document_ids同时派生 attention、loss 与 position 契约,但三者解决不同问题。- 同文档因果可见阻止信息泄漏,同文档标签转移阻止学习随机文档顺序。
- 最有力的单元测试是:关闭随机性后,每篇文档单独运行与装箱运行逐位置等价。
15 思考题与小练习#
- 文档长度
[5,4,3,2,2]、容量 8,用 First-Fit Decreasing 手算装箱结果、token 利用率与剩余空位。 - 修改代码,使“允许跨文档注意力、但不计算跨文档 loss”,解释它与完全隔离方案的数据分布差异。
- 为 packed batch 写三个断言:每个真实 query 至少看见自己、不能看未来、不同
document_id永不可见。
相关工作#
- Vaswani et al., Attention Is All You Need ↗,Transformer 与因果/填充注意力的基础。
- Raffel et al., Exploring the Limits of Transfer Learning with a Unified Text-to-Text Transformer ↗,讨论大规模文本预训练数据与序列构造。
- Krell et al., Efficient Sequence Packing without Cross-contamination ↗,系统研究无交叉污染的高效装箱。
- Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness ↗,解释高效 attention kernel 的 IO 约束。
- Kundu et al., Smart Batching: Fast Fine-Tuning of Transformer Language Models ↗,比较长度感知批处理与 padding 效率。
16 下一篇预告#
完全装箱适合可重排的预训练语料,但微调和在线任务常保留“一行一个样本”。下一篇将研究动态 padding、长度分桶与 token-based batching,回答怎样减少尾部浪费,又不改变样本权重和梯度尺度。