Skip to content

T5 span corruption and relative position buckets

ModuleL6.4 · side (optional, D31) · Python · Pass 5 · 2 to 3 h
You buildpython/tinyllm/obj/t5.py: noise_span_counts, random_spans_noise_mask, span_corrupt, t5_relative_bucket, T5RelativeBias
Contractcourse/contracts/py/tinyllm/obj/t5.pyi
Testscourse/tests/L6.4/test_t5.py (what they check: section 4); the bucket oracle is Hugging Face’s own T5 code
NeedsL0.1 Tensor · L0.2 permute · L0.4 Embedding · reading: L5.5 encoder-decoder Transformer, M07.1 sampling, M06.3 PCG32 (or --ref-deps)
Used byno module: the side quest sq.t5 trains an encoder-decoder with this objective
Milestonenone (optional module)
Optional depthRaffel et al., “Exploring the Limits of Transfer Learning with a Unified Text-to-Text Transformer” (2020), sections 2.1, 3.1.3, and 3.3; Shaw, Uszkoreit, and Vaswani, “Self-Attention with Relative Position Representations” (2018)
  • Span corruption hides about 15% of the tokens in a few whole spans, one sentinel id per span; inputs plus targets hold every original token exactly once (test_inputs_and_targets_restore_the_text).
  • The counts are fixed by the length, the density, and the mean span (round half to even), and the mask always starts with kept tokens and ends with a span (test_counts_and_lengths).
  • The random choices are two Fisher-Yates shuffles with one below draw per swap, noise lengths first, so a seed replays the batch (test_draws_follow_the_spec).
  • T5 has no position embeddings: each head adds a learned bias indexed by a bucket of the key’s offset, exact when small and logarithmic up to max_distance (test_buckets_match_hf).
  • A causal decoder folds every future offset into bucket 0 and uses all buckets for the past (test_hand_example_buckets, test_relative_bias_matches_hf).
Terminal window
ol start L6.4 # stubs t5.py
ol tests L6.4 # the course tests
ol check L6.4 # exit code is the verdict
ol diff L6.4 # after passing: your code against the reference

You have now met two pretraining objectives: GPT’s next token (L6.1) and BERT’s masked tokens (L6.2). BERT’s targets are single tokens at known positions, which an encoder-decoder (L5.5) cannot use well: its decoder writes sequences. T5 recast pretraining as text to text: corrupt the input by dropping whole spans and ask the decoder to write them back. It also replaced absolute position embeddings with a small learned table indexed by relative distance, which generalizes past the training length. Neither piece has a call site in your system (D31), which is why this module is optional; the side quest sq.t5 uses both to train a small T5 on your corpus.

SymbolMeaningType / shape
LLthe number of tokens of one textint ≥2\ge 2
ρ\rhonoise_density, the fraction to corruptfloat in (0,1)(0, 1)
μ\mumean_noise_span, the average span lengthfloat ≥1\ge 1
NNnoise tokens, clamp(round(Lρ),1,L−1)\mathrm{clamp}(\mathrm{round}(L \rho), 1, L - 1)int
KKnoise spans, clamp(round(N/μ),1,min⁡(N,L−N))\mathrm{clamp}(\mathrm{round}(N / \mu), 1, \min(N, L - N))int
SkS_ksentinel kk, the id sentinel_start_id −k- kint
δ=j−i\delta = j - irelative position of key jj from query iiint
nnbuckets per direction (num_buckets, halved when bidirectional)int
DDmax_distanceint

T5’s random_spans_noise_mask decides how many before where. With NN noise tokens in KK spans, it splits NN into KK positive span lengths and the L−NL - N kept tokens into KK positive gap lengths, then lays them out as keep, noise, keep, noise, and so on: the text always starts with kept tokens and ends with a noise span. Splitting mm items into kk positive parts uniformly is a shuffle: take m−1m - 1 flags, the first k−1k - 1 of them true, shuffle them, and start a new part after every true flag. The shuffle is Fisher-Yates from the end, $j = $ rng.below(i + 1) for i=m−2i = m - 2 down to 1; the noise lengths are drawn first.

Rounding is numpy’s (half to even): L=10L = 10, ρ=0.25\rho = 0.25 gives N=round(2.5)=2N = \mathrm{round}(2.5) = 2, not 3. T5 clamps KK only from below; on a short text round(N/μ)\mathrm{round}(N / \mu) can exceed L−NL - N, and then some gap would be empty, so the course clamps K≤min⁡(N,L−N)K \le \min(N, L - N).

Each noise span kk is replaced in the inputs by sentinel SkS_k, and the targets are each span preceded by its sentinel. T5’s sentinels are the last ids of the vocabulary counting down (<extra_id_0> is V−1V - 1), so sentinel_start_id is V−1V - 1 and Sk=V−1−kS_k = V - 1 - k. Real ids must stay out of the sentinel range, or the targets become ambiguous. Replacing every sentinel of the inputs by what follows it in the targets gives the text back: the objective loses nothing, and the targets are short (about N+KN + K tokens).

Self-attention without positions is permutation-equivariant. T5 adds to head hh‘s logit for query ii and key jj a learned scalar bh[bucket(δ)]b_h[\mathrm{bucket}(\delta)], with δ=j−i\delta = j - i. Small offsets matter most and get one bucket each; large offsets share logarithmically wider buckets:

bucket(r)={rr<n/2min⁡ ⁣(n−1, n/2+⌊ln⁡(r/(n/2))ln⁡(D/(n/2)) (n−n/2)⌋)otherwise\mathrm{bucket}(r) = \begin{cases} r & r < n/2 \\ \min\!\Big(n - 1,\ n/2 + \Big\lfloor \frac{\ln(r / (n/2))}{\ln(D / (n/2))}\,(n - n/2) \Big\rfloor\Big) & \text{otherwise} \end{cases}

for a distance r≥0r \ge 0. Bidirectional (the encoder): nn is half of num_buckets, r=∣δ∣r = \lvert\delta\rvert, and keys after the query (δ>0\delta > 0) add nn, so the two directions use disjoint halves. Causal (the decoder): nn is all of num_buckets, r=max⁡(0,−δ)r = \max(0, -\delta): every future key falls into bucket 0, which the causal mask hides anyway. The logarithm is computed in float32 and truncated toward zero, exactly as the Hugging Face code that T5 checkpoints were trained with; the test compares every offset in −300..300-300..300. The bias table is an Embedding(num_buckets, n_heads) named relative_attention_bias, and T5RelativeBias.forward(q_len, k_len, q_offset) returns [n_heads, q_len, k_len], with q_offset the tokens already in a decoder’s cache.

One span. Tokens 10 to 17 (L=8L = 8), ρ=0.25\rho = 0.25, μ=2\mu = 2: N=round(2)=2N = \mathrm{round}(2) = 2, K=round(1)=1K = \mathrm{round}(1) = 1. One span leaves no choice: 6 kept tokens, then 2 noise tokens. With S0=99S_0 = 99: inputs (10,11,12,13,14,15,99)(10, 11, 12, 13, 14, 15, 99), targets (99,16,17)(99, 16, 17).

Two spans, scripted draws. Tokens 0 to 9, ρ=0.4\rho = 0.4, μ=2\mu = 2: N=4N = 4, K=2K = 2.

stepflags beforedrawflags afterlengths
noise: 3 flags, first 1 trueT F Fbelow(3) = 0, swap 2 and 0F F T
F F Tbelow(2) = 1, swap 1 and 1F F Tparts after each true flag: 3, 1
keep: 5 flags, first 1 trueT F F F Fbelow(5) = 2, below(4) = 0, below(3) = 2, below(2) = 0F F F T F4, 2

Layout: keep 4, noise 3, keep 2, noise 1, so tokens 4, 5, 6 and 9 are noise. Inputs (0,1,2,3,S0,7,8,S1)(0, 1, 2, 3, S_0, 7, 8, S_1), targets (S0,4,5,6,S1,9)(S_0, 4, 5, 6, S_1, 9).

Buckets. Bidirectional, 32 buckets (n=16n = 16, exact below 8), D=128D = 128:

δ\deltarrbucket
−3-333
+3+3316+3=1916 + 3 = 19
−20-20208+⌊ln⁡2.5/ln⁡16×8⌋=8+⌊2.64⌋=108 + \lfloor \ln 2.5 / \ln 16 \times 8 \rfloor = 8 + \lfloor 2.64 \rfloor = 10
+200+20020016+min⁡(15,8+⌊9.29⌋)=3116 + \min(15, 8 + \lfloor 9.29 \rfloor) = 31
000

Causal with 32 buckets: δ=+5\delta = +5 gives bucket 0, δ=−5\delta = -5 gives 5.

These are test_hand_example_one_span, test_hand_example_two_spans_scripted_draws, and test_hand_example_buckets.

def noise_span_counts(length: int, noise_density: float, mean_noise_span: float) -> tuple[int, int]: ...
def random_spans_noise_mask(length, noise_density, mean_noise_span, rng) -> NDArray: ... # bool [L]
def span_corrupt(ids, noise_density, mean_noise_span, sentinel_start_id, rng, eos_id=None) -> tuple[NDArray, NDArray]: ...
def t5_relative_bucket(rel_pos, bidirectional: bool, num_buckets=32, max_distance=128) -> NDArray: ...
class T5RelativeBias(Module):
def __init__(self, n_heads, num_buckets=32, max_distance=128, bidirectional=True, rng=None): ...
def forward(self, q_len: int, k_len: int, q_offset: int = 0) -> Tensor: ... # [H, q, k]
TestKINDChecksWhy it matters downstream
test_hand_example_one_spanunitsection 3, with and without eos_idyou and the test agree on the layout
test_hand_example_two_spans_scripted_drawsunitthe six scripted draws, in order, give section 3’s inputs and targetsthe shuffle and the sentinel order
test_hand_example_bucketsunitsection 3’s bucket table, both directionsthe bucket formula
test_inputs_and_targets_restore_the_textproperty40 random texts: putting spans back gives the idsthe objective loses nothing
test_counts_and_lengthspropertyNN, KK (half to even), mask starts kept and ends noisy, sequence lengthsthe batch shapes a trainer pads to
test_short_texts_never_make_empty_spansboundarythe upper clamp of KK on 10 and 6 tokensno empty span or gap
test_draws_follow_the_specunitthe exact below calls; same seed, same outputreproducible data
test_sentinels_and_validationboundarysentinels count down; colliding ids, no room, bad densities, L<2L < 2 raiseunambiguous targets
test_buckets_match_hfgoldenfive settings over −300..300-300..300 against HF’s _relative_position_bucketpretrained T5 tables index correctly
test_bucket_propertiespropertymonotone, in range, disjoint halves, exact while smallthe coarsening is a coarsening
test_relative_bias_matches_hfgoldenHF compute_bias, encoder and decoder with past tokensthe bias each head adds
test_relative_bias_gradient_lands_on_bucketspropertythe gradient is a scatter-add of the upstream gradient by buckettraining the table
PitfallSymptomCaught by
1. relative position as query minus key, or a causal decoder that gives future keys their distancepast and future swapped: pretrained biases read the wrong sidetest_relative_bias_matches_hf (mutant s01), test_buckets_match_hf (mutant s09)
2. bidirectional buckets not halvedthe two directions overlap; HF disagrees past offset 8test_buckets_match_hf, test_bucket_properties (mutant s02)
3. rounding half up; no upper clamp on the span count3 noise tokens where T5 has 2; empty spans on short textstest_counts_and_lengths (mutant s03), test_short_texts_never_make_empty_spans (mutant s10)
4. sentinels counting up, or targets without themids collide with real tokens; the decoder cannot tell spans aparttest_hand_example_two_spans_scripted_draws, test_sentinels_and_validation (mutants s04, s05)
5. a text that starts with a noise spanthe first tokens always corrupted, unlike T5test_hand_example_one_span, test_counts_and_lengths (mutant s06)
6. the shuffle’s range or draw order changeda seed no longer replays the batchtest_draws_follow_the_spec (mutants s07, s08)
DirectionModuleHow it uses this
BackL0.1the bias is a Tensor whose gradient trains the table
BackL0.2permute turns the [q, k, H] lookup into [H, q, k]
BackL0.4the table is an Embedding(num_buckets, n_heads)
BackL5.5the encoder-decoder that span corruption trains (reading)
Forwardsq.t5a small T5 on your corpus: span corruption for pretraining, the bias inside every attention layer

No module calls this code (D31): L6.4 is optional. It is still worth an evening: span corruption is the objective of T5, UL2, and many code models, and relative buckets are the simplest relative position scheme, a good contrast to RoPE (L7.3) and ALiBi (L7.4).

Your pieceProduction equivalentWhat it addsWhere to look
span_corruptT5’s span_corruption preprocessor; HF DataCollatorForT5MLMpacking to fixed input and target lengths, EOS handling, expand_inputs_and_targets to choose the raw lengtht5/data/preprocessors.py; transformers/examples/flax/language-modeling/run_t5_mlm_flax.py
the denoising mixtureUL2’s mixture of denoisersshort spans, long spans, and prefix LM in one model, chosen by a mode tokenTay et al., “UL2: Unifying Language Learning Paradigms” (2022)
T5RelativeBiasHF T5Attention.compute_biasthe bias computed once in layer 0 and reused by every layertransformers/models/t5/modeling_t5.py
relative positionsRoPE, ALiBiposition by rotation, or a fixed linear penalty with no tableL7.3, L7.4