Course Content
Transformer Architecture Q&A
6 sections · 60 lessons
When training with packed documents, what issues do absolute positional embeddings introduce?
What you need to know
If your data has many 300-token documents and your sequence length is 8,192, padding each document would waste over 95% of compute. Packing fixes the waste but creates two problems.
Problem 1: position offset
With naive packing, position ids run 0 to 8,191 across the whole packed sequence. A document that starts at slot 3,000 has its first token labelled position 3,000. With learned absolute embeddings, positions 0 and 3,000 are unrelated vectors, so the model must learn "start of document" at every possible offset. At inference, documents always start at position 0, so training and inference differ.
Problem 2: cross-document attention
With only a causal mask, a token in document B can attend to every token of document A before it. The model learns links between unrelated texts — for example, it learns that a bank statement can predict the next word of a recipe. It also lets content from one sample leak into another.
The fix, worked by hand
Three documents of lengths 3, 2 and 3 packed into 8 slots:
1import torch23doc_lens = [3, 2, 3] # three documents packed into 8 slots4doc_id = torch.repeat_interleave(torch.arange(3), torch.tensor(doc_lens))5position_ids = torch.cat([torch.arange(n) for n in doc_lens])67same_doc = doc_id[:, None] == doc_id[None, :] # block diagonal8causal = torch.ones(8, 8, dtype=torch.bool).tril()9allowed = same_doc & causal # True = may attend1011print(doc_id.tolist()) # [0, 0, 0, 1, 1, 2, 2, 2]12print(position_ids.tolist()) # [0, 1, 2, 0, 1, 0, 1, 2]13print(allowed.int())allowed (1 = may attend)1 0 0 0 0 0 0 01 1 0 0 0 0 0 01 1 1 0 0 0 0 00 0 0 1 0 0 0 00 0 0 1 1 0 0 00 0 0 0 0 1 0 00 0 0 0 0 1 1 00 0 0 0 0 1 1 1Positions restart at 0 for each document, and the mask is three small lower triangles. In production you do not build this 8 × 8 (or 8,192 × 8,192) matrix. You pass cumulative sequence lengths (cu_seqlens = [0, 3, 5, 8]) to a variable-length kernel such as FlashAttention's flash_attn_varlen_func, or use a document mask in PyTorch's FlexAttention. Hugging Face's DataCollatorWithFlattening builds the reset position ids for you.
Does RoPE fix it?
Only the first problem, and only partly: RoPE scores depend on distance, so an offset matters less. The second problem remains with any position scheme — you still need the document mask.
A real-life example
A team pre-trains a small code-completion model on millions of Python files, most under 500 tokens, packed into 4,096-token sequences. Evaluation on single files looks fine, but in production the model sometimes suggests imports from libraries that the current file never uses.
Inspection shows the packer used a plain causal mask, so each file could attend to the tail of whatever unrelated file came before it. The model had learned to continue "whatever code came earlier", including other projects' imports. Switching to per-file position ids and variable-length attention removed the pattern, at no extra compute cost.
Follow-up questions to expect
- "Why not just pad?" — Padding wastes compute on tokens that carry no signal; with many short documents the waste is most of the batch.
- "Is cross-document attention ever acceptable?" — Some pre-training runs accept it for simplicity, with a separator token between documents, but a document mask is cleaner and now cheap.
- "What else must reset at boundaries?" — The loss should not ask the last token of one document to predict the first token of the next.