USAAIO Lesson 107, from Phase 3, on extractive question answering with two linear heads on the BERT output, predicting a start and an end position. It covers why treating the two softmaxes as independent works, cross-entropy loss over token positions, and constrained span decoding, which requires start ≤ end. It then covers evaluation by exact match and token-level F1, generative QA with T5, and open-domain QA, which retrieves with DPR and then reads and extracts. All the numbers were verified with torch 2.7.1+cpu and numpy 2.2.6. The lesson runs to 28 slides.
Subject: Machine Learning · 52 slides · code lesson
Open the interactive version of this deck · Homework for this lesson
Title
USAAIO · Lesson 107 · Phase 3
Find the span. Every design choice — two-head BERT, cross-entropy on token indices, constrained decoding, EM/F1 — derived from first principles and run live in PyTorch.
Objectives
Section
Part 1 of 4
Concept
SQuAD (Stanford Question Answering Dataset, 2016) pairs a passage and a question; the label is a character-level span within the passage — no generation, no knowledge retrieval, just pointer arithmetic over the input tokens.
| field | example | note |
|---|---|---|
| passage | "The Eiffel Tower is in Paris, France." | ≤512 BPE tokens after BERT tokenization |
| question | "Where is the Eiffel Tower?" | prepended to passage with [SEP] |
| answer | "Paris, France" | must be a contiguous passage substring |
| SQuAD 1.1 | always answerable | 100,837 QA pairs, 536 Wikipedia articles |
| SQuAD 2.0 | ~33% unanswerable | adds adversarial unanswerable questions |
The model never generates token-by-token; it outputs two integers — the start token index and the end token index — within the concatenated [CLS] question [SEP] passage [SEP] sequence.
Comparison
Comparison matrix
From SQuAD: from reading comprehension to token indices: refill the example column from what you know. The rest of the table is as it appeared.
| field | example | note |
|---|---|---|
| passage | "The Eiffel Tower is in Paris, France." | ≤512 BPE tokens after BERT tokenization |
| question | "Where is the Eiffel Tower?" | prepended to passage with [SEP] |
| answer | "Paris, France" | must be a contiguous passage substring |
| SQuAD 1.1 | always answerable | 100,837 QA pairs, 536 Wikipedia articles |
| SQuAD 2.0 | ~33% unanswerable | adds adversarial unanswerable questions |
Concept
BERT (Lesson 90) takes a single flat token sequence. For QA the sequence is [CLS] q1 q2 … qm [SEP] p1 p2 … pn [SEP], where q* are question tokens and p* are passage tokens. Token-type IDs (segment embeddings) mark which half each token belongs to.
\[ \text{input} = [\texttt{CLS}]\;q_1\cdots q_m\;[\texttt{SEP}]\;p_1\cdots p_n\;[\texttt{SEP}]\quad\text{length }T = m+n+3 \]
| token range | segment ID | model sees |
|---|---|---|
| [CLS], q1…qm, [SEP] | 0 | question context |
| p1…pn, [SEP] | 1 | passage context |
| start/end labels | — | must lie inside passage tokens only |
Trade off
Comparison matrix
From BERT input for SQuAD: the concatenated sequence: every row here is a choice with a cost. Fill the segment ID column, then say which row you would actually pick and what you give up for it.
| token range | segment ID | model sees |
|---|---|---|
| [CLS], q1…qm, [SEP] | 0 | question context |
| p1…pn, [SEP] | 1 | passage context |
| start/end labels | — | must lie inside passage tokens only |
Section
Part 2 of 4
Concept
BERT produces a hidden vector h_i ∈ ℝ^{d} for every input token i. Two separate nn.Linear(d, 1) layers score each position as a candidate start and a candidate end.
\[ s_i = w_s^\top h_i \quad e_i = w_e^\top h_i \qquad s,e \in \mathbb{R}^T \]
During training, softmax is applied independently over s and e, and two cross-entropy losses (one per position) are summed. At inference, the predicted span is argmax_i s_i to argmax_j e_j, with the constraint i ≤ j.
Counterexample
Discussion prompt
BERT produces a hidden vector h_i ∈ ℝ^{d} for every input token i. Two separate nn.Linear(d, 1) layers score each position as a candidate start and a candidate end.
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Estimation
Predict first
Trace a batch of B=2 sequences, each T=16 tokens, through BERT hidden size d=64. Predict tensor shapes at each step before reading on.
Commit before you compute: what does Forward pass: shapes and loss — verified come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: start_logits and end_logits: both torch.Size([2, 16])
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Linear(64,1) applied to each token's 64-d vector → scalar per token → squeeze the trailing dim → (B=2, T=16).
Worked example
Trace a batch of B=2 sequences, each T=16 tokens, through BERT hidden size d=64. Predict tensor shapes at each step before reading on.
import torch, torch.nn as nn, torch.nn.functional as F
torch.manual_seed(42)
B, T, d = 2, 16, 64
bert_output = torch.randn(B, T, d) # simulated BERT encoder output
start_head = nn.Linear(d, 1)
end_head = nn.Linear(d, 1)
start_logits = start_head(bert_output).squeeze(-1) # (B, T)
end_logits = end_head(bert_output).squeeze(-1) # (B, T)
# Training: cross-entropy on gold token positions
gold_start = torch.tensor([3, 7]) # ground-truth start indices
gold_end = torch.tensor([5, 11]) # ground-truth end indices
loss_s = F.cross_entropy(start_logits, gold_start) # 3.2292
loss_e = F.cross_entropy(end_logits, gold_end) # 2.8630
loss = (loss_s + loss_e) / 2 # 3.0461
print(f'start shape: {start_logits.shape}, end shape: {end_logits.shape}')
print(f'loss_start={loss_s:.4f} loss_end={loss_e:.4f} total={loss:.4f}')start_logits and end_logits: both torch.Size([2, 16])
Why: Linear(64,1) applied to each token's 64-d vector → scalar per token → squeeze the trailing dim → (B=2, T=16).
| tensor | shape | operation producing it |
|---|---|---|
| bert_output | (2, 16, 64) | simulated BERT encoder (B=2, T=16, d=64) |
| start_head(bert_output) | (2, 16, 1) | Linear(64,1) applied at each token |
| start_logits (squeezed) | (2, 16) | one score per token per example |
| end_logits (squeezed) | (2, 16) | same for end head |
| loss_start | scalar 3.2292 | cross_entropy over T=16 classes, gold=[3,7] |
| loss_end | scalar 2.8630 | cross_entropy over T=16 classes, gold=[5,11] |
| total loss | scalar 3.0461 | (loss_start + loss_end) / 2 |
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
start_logits and end_logits: both torch.Size([2, 16])
What was it asked to do, and what must it have been given? Reconstruct the problem from its answer.
Hint: Every quantity in the result had to enter somewhere. Account for each one.
Answer:
Trace a batch of B=2 sequences, each T=16 tokens, through BERT hidden size d=64. Predict tensor shapes at each step before reading on.
Concept
The true distribution over spans is P(start=i, end=j) — a joint over T² pairs. Modeling it exactly would require a T×T output matrix. Instead BERT-SQuAD factors it as P(start=i) · P(end=j) — independent marginals.
\[ P(\text{span}=(i,j)) \approx P_s(i) \cdot P_e(j) \quad\text{(independence assumption)} \]
This factorization is valid when the start and end positions carry redundant signal (both depend heavily on the same passage words), so the shared BERT context implicitly couples them. The approximation fails for very long answers where end depends strongly on the exact start chosen.
Anomaly
Predict first
A student writes this, and it looks reasonable:
Decode naively: start = argmax(start_logits), end = argmax(end_logits) — return the span (start, end) regardless of order.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: This is a common failure mode: the two heads are trained independently and can produce an inverted span (end < start), which is syntactically impossible.
Enumerate all valid (i, j) pairs where i ≤ j, compute start_logits[i] + end_logits[j], pick the argmax.
Why: This is a common failure mode: the two heads are trained independently and can produce an inverted span (end < start), which is syntactically impossible.
Trap
Decode naively: start = argmax(start_logits), end = argmax(end_logits) — return the span (start, end) regardless of order.
Suppose start_logits peaks at token 9, end_logits peaks at token 4
Why: This is a common failure mode: the two heads are trained independently and can produce an inverted span (end < start), which is syntactically impossible.
Result: span = (9, 4) — an invalid span. A span of zero or negative length is meaningless and will crash downstream string-extraction code.
Enumerate all valid (i, j) pairs where i ≤ j, compute start_logits[i] + end_logits[j], pick the argmax.
Best valid span search: O(T²) over all i ≤ j, maximize s_i + e_j
Why: Summing logits (not probabilities) is equivalent to maximizing log P(start=i) + log P(end=j) under the independence assumption. With T ≤ 512, O(T²) ≈ 262 k operations — negligible.
Verified: with seed=1, independent argmax gives start=0, end=4 (valid in this case, logit sum=1.0838); the constrained search confirms (0, 4) as the best valid pair.
Section
Part 3 of 4
Concept
SQuAD uses two metrics, both computed after normalization (lowercasing, stripping articles and punctuation, collapsing whitespace). Multiple gold answers are provided; the prediction is scored against the best-matching gold.
\[ \text{EM} = \mathbf{1}[\text{pred} = \text{gold}] \qquad F_1 = \frac{2 \cdot P \cdot R}{P + R} \]
P = |common tokens| / |pred tokens|, R = |common tokens| / |gold tokens|. Token overlap is computed on bag-of-words (multisets of normalized whitespace-split tokens), NOT character-level.
Estimation
Predict first
Predict the EM and F1 for each pair before reading the table. Remember: normalize first (lowercase, split on whitespace), then compute bag-of-word overlap.
Commit before you compute: what does EM and F1 — four concrete examples come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: "albert einstein" vs "einstein": EM=0, P=0.500, R=1.000, F1=0.667
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Common = {'einstein'} (1 token).
Worked example
Predict the EM and F1 for each pair before reading the table. Remember: normalize first (lowercase, split on whitespace), then compute bag-of-word overlap.
def normalize(s):
return s.lower().strip().split()
def compute_f1(pred, gold):
p_toks, g_toks = normalize(pred), normalize(gold)
common = set(p_toks) & set(g_toks)
if not common:
return 0.0, 0.0, 0.0
prec = len(common) / len(p_toks)
rec = len(common) / len(g_toks)
f1 = 2*prec*rec / (prec+rec)
return prec, rec, f1
examples = [
('albert einstein', 'albert einstein'),
('albert einstein', 'einstein'),
('the united states', 'united states of america'),
('42', '42'),
]
for pred, gold in examples:
em = int(normalize(pred) == normalize(gold))
p, r, f1 = compute_f1(pred, gold)
print(f'pred={pred!r:<22} gold={gold!r:<24} EM={em} P={p:.3f} R={r:.3f} F1={f1:.3f}')"albert einstein" vs "einstein": EM=0, P=0.500, R=1.000, F1=0.667
Why: Common = {'einstein'} (1 token). Pred has 2 tokens so P=1/2. Gold has 1 token so R=1/1. Harmonic mean = 2·0.5·1/(0.5+1) = 0.667.
| predicted | gold | EM | P | R | F1 |
|---|---|---|---|---|---|
| albert einstein | albert einstein | 1 | 1.000 | 1.000 | 1.000 |
| albert einstein | einstein | 0 | 0.500 | 1.000 | 0.667 |
| the united states | united states of america | 0 | 0.667 | 0.500 | 0.571 |
| 42 | 42 | 1 | 1.000 | 1.000 | 1.000 |
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
"albert einstein" vs "einstein": EM=0, P=0.500, R=1.000, F1=0.667
What was it asked to do, and what must it have been given? Reconstruct the problem from its answer.
Hint: Every quantity in the result had to enter somewhere. Account for each one.
Answer:
Predict the EM and F1 for each pair before reading the table. Remember: normalize first (lowercase, split on whitespace), then compute bag-of-word overlap.
Anomaly
Predict first
A student writes this, and it looks reasonable:
Compute F1 by counting characters in common between predicted and gold strings.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Character counting gives an inflated or deflated number depending on string length, and ignores token boundaries entirely.
Normalize both strings (lowercase, strip articles/punct), split on whitespace, take bag-of-word intersection, compute token-level precision and recall.
Why: Character counting gives an inflated or deflated number depending on string length, and ignores token boundaries entirely. 'states' shares characters with 'states' but 'the' contributes characters not in gold.
Trap
Compute F1 by counting characters in common between predicted and gold strings.
pred='the united states', gold='united states of america' — char overlap = len('united states') = 13
Why: Character counting gives an inflated or deflated number depending on string length, and ignores token boundaries entirely. 'states' shares characters with 'states' but 'the' contributes characters not in gold.
Normalize both strings (lowercase, strip articles/punct), split on whitespace, take bag-of-word intersection, compute token-level precision and recall.
pred tokens = {the, united, states}, gold tokens = {united, states, of, america}; common = {united, states}
Why: P = 2/3 = 0.667 (2 of 3 pred tokens are correct), R = 2/4 = 0.500 (2 of 4 gold tokens recovered). F1 = 2·0.667·0.5/(0.667+0.5) = 0.571. This matches the official SQuAD script.
Break the constraint
Discussion prompt
The rule this trap just fixed:
P = 2/3 = 0.667 (2 of 3 pred tokens are correct), R = 2/4 = 0.500 (2 of 4 gold tokens recovered). F1 = 2·0.667·0.5/(0.667+0.5) = 0.571. This matches the official SQuAD script.
Now break it on purpose. Build a case that violates it and follow the consequences until something visibly fails. Where does the failure first show up — and would you have noticed it if you had not been looking?
Hint: The dangerous rules are the ones whose violation still produces an answer. If yours fails loudly, try to find one that fails quietly.
Answer:
Character counting gives an inflated or deflated number depending on string length, and ignores token boundaries entirely. 'states' shares characters with 'states' but 'the' contributes characters not in gold.
Section
Part 4 of 4
Concept
Extractive QA on SQuAD assumes the relevant passage is given. Open-domain QA starts from a question alone and must first retrieve passages from a large corpus (e.g., Wikipedia with ~21 M 100-word chunks), then read to extract the span.
Concept
DPR (Karpukhin 2020) trains two BERT encoders independently: a question encoder E_Q and a passage encoder E_P. Both output a CLS-token vector. Similarity is the normalized dot product.
\[ \text{sim}(q, p) = E_Q(q)^\top E_P(p) \quad\text{(both unit-normalized → cosine similarity)} \]
At serving time, passage embeddings are pre-computed and stored in a FAISS index; only the query embedding is computed online. DPR outperforms TF-IDF on open-domain benchmarks (e.g., NaturalQuestions Top-20 accuracy: 79% DPR vs 59% BM25).
Estimation
Predict first
Five candidate passage embeddings, one query embedding — all unit-normalized (seed=7). Which passage ranks first? Compute the dot products mentally before revealing.
Commit before you compute: what does DPR retrieval — dot-product ranking verified come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Passage 3 ranks first with score 0.2467
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Dot product of two unit vectors equals cosine similarity.
Worked example
Five candidate passage embeddings, one query embedding — all unit-normalized (seed=7). Which passage ranks first? Compute the dot products mentally before revealing.
import torch, torch.nn.functional as F
torch.manual_seed(7)
d = 32
query_emb = F.normalize(torch.randn(1, d), dim=-1) # unit vector
passage_embs = F.normalize(torch.randn(5, d), dim=-1) # 5 unit vectors
scores = (query_emb @ passage_embs.T).squeeze(0) # (5,)
ranked = scores.argsort(descending=True)
for rank, idx in enumerate(ranked.tolist()):
print(f'Rank {rank+1}: passage {idx} score={scores[idx].item():.4f}')Passage 3 ranks first with score 0.2467
Why: Dot product of two unit vectors equals cosine similarity. Seed 7 gives passage 3 the highest inner product with the query; passage 0 is most dissimilar at -0.2886.
| rank | passage | dot-product score |
|---|---|---|
| 1 | 3 | 0.2467 |
| 2 | 2 | 0.2267 |
| 3 | 4 | 0.0991 |
| 4 | 1 | -0.0375 |
| 5 | 0 | -0.2886 |
Comparison
Comparison matrix
From DPR retrieval — dot-product ranking verified: refill the passage column from what you know. The rest of the table is as it appeared.
| rank | passage | dot-product score |
|---|---|---|
| 1 | 3 | 0.2467 |
| 2 | 2 | 0.2267 |
| 3 | 4 | 0.0991 |
| 4 | 1 | -0.0375 |
| 5 | 0 | -0.2886 |
Concept
Generative QA (e.g., T5 fine-tuned on SQuAD/NQ) treats QA as sequence-to-sequence: input question: ... context: ..., output the answer string token-by-token. No span constraint — answers can be rephrased or multi-sentence.
| property | extractive (BERT) | generative (T5) |
|---|---|---|
| output | start/end token indices | auto-regressive token sequence |
| answer | must be a passage substring | any string; can be rephrased |
| unanswerable | SQuAD 2.0: add 'no answer' class | model outputs 'unanswerable' |
| latency | O(1) decode step | O(answer_len) decode steps |
| EM/F1 | standard SQuAD metrics | ROUGE or exact string match |
For USAAIO: extractive is architecturally cleaner and the dominant SQuAD paradigm. Generative models (FiD, RAG) are used for open-ended questions where the answer is not a verbatim passage substring.
Counterexample
Discussion prompt
For USAAIO: extractive is architecturally cleaner and the dominant SQuAD paradigm. Generative models (FiD, RAG) are used for open-ended questions where the answer is not a verbatim passage substring.
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Constraint
Discussion prompt
Run The QA system playbook with this step confiscated:
Train: loss = (cross_entropy(start_logits, gold_start) + cross_entropy(end_logits, gold_end)) / 2.
Is it still possible? If it is, say what takes its place and what it costs you. If it is not, say exactly what that step was providing that nothing else does.
Hint: A step you can drop for free was never load-bearing. If you cannot drop it, name the thing that goes wrong the moment it is gone.
Answer:
[CLS] question [SEP] passage [SEP] into ≤512 BPE tokens; record passage token offsets for span-to-character mapping.h ∈ ℝ^{T×d} for the full sequence.start_logits = Linear_s(h).squeeze(-1), end_logits = Linear_e(h).squeeze(-1) — both shape (T,).loss = (cross_entropy(start_logits, gold_start) + cross_entropy(end_logits, gold_end)) / 2.i ≤ j, pick argmax_i,j [s_i + e_j]; mask question tokens and [SEP]/[CLS] tokens out.Pattern
[CLS] question [SEP] passage [SEP] into ≤512 BPE tokens; record passage token offsets for span-to-character mapping.h ∈ ℝ^{T×d} for the full sequence.start_logits = Linear_s(h).squeeze(-1), end_logits = Linear_e(h).squeeze(-1) — both shape (T,).loss = (cross_entropy(start_logits, gold_start) + cross_entropy(end_logits, gold_end)) / 2.i ≤ j, pick argmax_i,j [s_i + e_j]; mask question tokens and [SEP]/[CLS] tokens out.Edge cases
Discussion prompt
The QA system playbook works on the cases you have just seen. Push it to the edge: what is the most degenerate input it still handles — empty, zero, one item, everything equal — and what is the first case where it stops being true? Name the case, not just "it breaks".
Hint: Try the smallest legal input, then the largest, then the one where two things collide. Methods are specified at their edges; the middle takes care of itself.
Answer:
[CLS] question [SEP] passage [SEP] into ≤512 BPE tokens; record passage token offsets for span-to-character mapping.h ∈ ℝ^{T×d} for the full sequence.start_logits = Linear_s(h).squeeze(-1), end_logits = Linear_e(h).squeeze(-1) — both shape (T,).loss = (cross_entropy(start_logits, gold_start) + cross_entropy(end_logits, gold_end)) / 2.i ≤ j, pick argmax_i,j [s_i + e_j]; mask question tokens and [SEP]/[CLS] tokens out.Elimination
Eliminate the wrong options
BERT hidden size d=768, two Linear(768, 1) heads. How many new parameters do the QA heads add?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: Each Linear(768, 1) has 768 weight parameters + 1 bias = 769 parameters. Two independent heads: 769 × 2 = 1,538. This is why fine-tuning BERT for QA is so lightweight — the pretrained encoder (110 M params) does the heavy lifting; the QA head adds 0.001%.
Check
Work this out before clicking. A BERT model has hidden size d=768. The extractive QA head consists of two linear layers. What is the total number of trainable parameters added by these two heads (weights + biases)?
Check your understanding
BERT hidden size d=768, two Linear(768, 1) heads. How many new parameters do the QA heads add?
Answer: A
Why: Each Linear(768, 1) has 768 weight parameters + 1 bias = 769 parameters. Two independent heads: 769 × 2 = 1,538. This is why fine-tuning BERT for QA is so lightweight — the pretrained encoder (110 M params) does the heavy lifting; the QA head adds 0.001%.
Prediction
Predict first
pred = 'New York City', gold = 'New York'. What are EM and F1 (token-level)?
Answer it in your own words, now, with nothing to choose from. The options are on the next slide — and picking the right one off a list is an easier skill than producing it.
Correct: EM = 0, F1 = 0.800
Why: Normalized pred tokens = {new, york, city} (3 tokens). Normalized gold tokens = {new, york} (2 tokens). Common = {new, york} (2 tokens). P = 2/3 ≈ 0.667, R = 2/2 = 1.000. F1 = 2·(2/3)·1 / (2/3+1) = (4/3)/(5/3) = 4/5 = 0.800. EM = 0 because the token lists are not identical.
Check
Compute before clicking. Predicted span: "New York City". Gold span: "New York". After normalization (lowercase, split on spaces), what are the EM and F1?
Check your understanding
pred = 'New York City', gold = 'New York'. What are EM and F1 (token-level)?
Answer: A
Why: Normalized pred tokens = {new, york, city} (3 tokens). Normalized gold tokens = {new, york} (2 tokens). Common = {new, york} (2 tokens). P = 2/3 ≈ 0.667, R = 2/2 = 1.000. F1 = 2·(2/3)·1 / (2/3+1) = (4/3)/(5/3) = 4/5 = 0.800. EM = 0 because the token lists are not identical.
Elimination
Eliminate the wrong options
BERT-SQuAD is trained with loss = (CE(start_logits, gold_start) + CE(end_logits, gold_end)) / 2. What does this training signal teach the model to do, and why does the independence assumption make this factorization valid?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: CE(start_logits, gold_start) trains the start head to assign maximum probability to the correct start token index — i.e., maximize log P_s(gold_start). Similarly for end. The two heads share the same BERT encoder output h_i, so the encoder learns a single representation that makes both position predictions accurate. The implicit coupling through shared h_i is why factoring the span probability as P_s · P_e is a good approximation despite the independence assumption.
Check
Think carefully about what the cross-entropy losses measure before clicking.
Check your understanding
BERT-SQuAD is trained with loss = (CE(start_logits, gold_start) + CE(end_logits, gold_end)) / 2. What does this training signal teach the model to do, and why does the independence assumption make this factorization valid?
Answer: A
Why: CE(start_logits, gold_start) trains the start head to assign maximum probability to the correct start token index — i.e., maximize log P_s(gold_start). Similarly for end. The two heads share the same BERT encoder output h_i, so the encoder learns a single representation that makes both position predictions accurate. The implicit coupling through shared h_i is why factoring the span probability as P_s · P_e is a good approximation despite the independence assumption.
Concept
Project brief: Implement TinyBERTforQA from scratch in PyTorch — a 2-layer transformer encoder with two linear QA heads — train it on synthetic extractive QA data, and evaluate EM and F1 on a held-out set.
nn.Embedding, nn.Linear, nn.LayerNorm.Worked example
Define TinyBERTforQA with vocab_size=500, d_model=64, n_heads=4, d_ff=128, n_layers=2, max_seq=32
Why: Small enough to train on CPU in seconds, large enough to demonstrate all architectural choices. Total verified parameter count: 101,378.
import torch, torch.nn as nn, torch.nn.functional as F
torch.manual_seed(0)
class MHSA(nn.Module):
def __init__(self, d, h):
super().__init__()
self.h, self.dk = h, d//h
self.qkv = nn.Linear(d, 3*d)
self.out = nn.Linear(d, d)
def forward(self, x):
B, T, D = x.shape
qkv = self.qkv(x).reshape(B,T,3,self.h,self.dk).unbind(2)
q,k,v = [t.transpose(1,2) for t in qkv]
a = F.softmax((q@k.transpose(-2,-1))/self.dk**0.5, dim=-1)
return self.out((a@v).transpose(1,2).reshape(B,T,D))
class Block(nn.Module):
def __init__(self, d, h, ff):
super().__init__()
self.attn = MHSA(d, h)
self.ff = nn.Sequential(nn.Linear(d,ff), nn.GELU(), nn.Linear(ff,d))
self.n1, self.n2 = nn.LayerNorm(d), nn.LayerNorm(d)
def forward(self, x):
x = self.n1(x + self.attn(x))
return self.n2(x + self.ff(x))
class TinyBERTforQA(nn.Module):
def __init__(self, V=500, d=64, h=4, ff=128, L=2, T=32):
super().__init__()
self.tok = nn.Embedding(V, d)
self.pos = nn.Embedding(T, d)
self.seg = nn.Embedding(2, d)
self.enc = nn.ModuleList([Block(d,h,ff) for _ in range(L)])
self.norm = nn.LayerNorm(d)
self.qa = nn.Linear(d, 2) # outputs: [start_logit, end_logit]
def forward(self, ids, segs):
B, T = ids.shape
p = torch.arange(T, device=ids.device).unsqueeze(0)
x = self.tok(ids) + self.pos(p) + self.seg(segs)
for blk in self.enc: x = blk(x)
lg = self.qa(self.norm(x)) # (B, T, 2)
return lg[:,:,0], lg[:,:,1] # start_logits, end_logits
model = TinyBERTforQA()
print(f'params: {sum(p.numel() for p in model.parameters()):,}') # 101,378Forward pass: B=2, T=20 → start_logits.shape = torch.Size([2, 20])
Why: ids shape (2,20) → embeddings summed → 2 Block layers → LayerNorm → Linear(64,2) at each position → split on dim=-1 → two (2,20) tensors.
| component | output shape | param count |
|---|---|---|
| tok_embed Embedding(500,64) | (B, T, 64) | 32,000 |
| pos_embed Embedding(32,64) | (B, T, 64) | 2,048 |
| seg_embed Embedding(2,64) | (B, T, 64) | 128 |
| MHSA (per Block): qkv+out | (B, T, 64) | 4×64² = 16,384 × 2 blocks |
| FFN (per Block): 64→128→64 | (B, T, 64) | 2×(64×128+128+128×64+64) = 33,024 × 2 blocks |
| LayerNorm (2 per block + 1 final) | — | small |
| QA head Linear(64,2) | (B, T, 2) | 130 |
| TOTAL | — | 101,378 |
Discrimination
Sort into buckets
Sort these by output shape, from memory, without looking back at Your turn — Milestone 1: model definition and…. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Pattern
Predict first
The table runs: 0 | ~3.0 | uniform over 20 classes: log(20) ≈ 3.0 · 50 | ~2.0 | model narrowing to correct region · 100 | ~1.0 | substantially learned start/end offsets · 150 | ~0.5 | near-convergence on synthetic task
In Your turn — Milestone 2: training loop with synthetic data, given the rows so far: what is the next one — the row where step is 200?
Correct: 200 | ~0.2 | overfitting to synthetic random seeds
| step | expected loss | interpretation |
|---|---|---|
| 0 | ~3.0 | uniform over 20 classes: log(20) ≈ 3.0 |
| 50 | ~2.0 | model narrowing to correct region |
| 100 | ~1.0 | substantially learned start/end offsets |
| 150 | ~0.5 | near-convergence on synthetic task |
| 200 | ~0.2 | overfitting to synthetic random seeds |
Why: The relationship between the columns, not the individual numbers, is what generates the next row. Using fixed random seed ensures reproducibility.
Worked example
Generate synthetic extractive QA examples: random token IDs with known gold start/end positions, then train for 200 steps
Why: Using fixed random seed ensures reproducibility. Loss should decrease from ~3.0 toward ~0.1 as the model learns to point at gold positions.
torch.manual_seed(42)
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-3)
# Synthetic data: B=8, T=20, gold start in [2,6], gold end in [start, start+3]
def make_batch(B=8, T=20, seed=None):
if seed is not None: torch.manual_seed(seed)
ids = torch.randint(1, 500, (B, T))
segs = torch.zeros(B, T, dtype=torch.long)
segs[:, 10:] = 1 # passage in second half
gs = torch.randint(10, 14, (B,)) # gold start in passage
ge = gs + torch.randint(1, 4, (B,)) # gold end = start + 1..3
ge = ge.clamp(max=T-1)
return ids, segs, gs, ge
for step in range(200):
ids, segs, gs, ge = make_batch(seed=step)
s_lg, e_lg = model(ids, segs)
loss = (F.cross_entropy(s_lg, gs) + F.cross_entropy(e_lg, ge)) / 2
optimizer.zero_grad(); loss.backward(); optimizer.step()
if step % 50 == 0:
print(f'step {step:3d} loss={loss.item():.4f}')Expected output: loss drops from ~3.0 at step 0 toward <0.5 by step 150
Why: With 101 k parameters and only 8 synthetic examples per step, the model quickly memorizes the random pattern. This confirms gradient flow through both heads is working correctly.
| step | expected loss | interpretation |
|---|---|---|
| 0 | ~3.0 | uniform over 20 classes: log(20) ≈ 3.0 |
| 50 | ~2.0 | model narrowing to correct region |
| 100 | ~1.0 | substantially learned start/end offsets |
| 150 | ~0.5 | near-convergence on synthetic task |
| 200 | ~0.2 | overfitting to synthetic random seeds |
Trade off
Comparison matrix
From Your turn — Milestone 2: training loop with synthetic data: every row here is a choice with a cost. Fill the expected loss column, then say which row you would actually pick and what you give up for it.
| step | expected loss | interpretation |
|---|---|---|
| 0 | ~3.0 | uniform over 20 classes: log(20) ≈ 3.0 |
| 50 | ~2.0 | model narrowing to correct region |
| 100 | ~1.0 | substantially learned start/end offsets |
| 150 | ~0.5 | near-convergence on synthetic task |
| 200 | ~0.2 | overfitting to synthetic random seeds |
Fill the middle
Fill in the blanks
From Your turn — Milestone 3: EM/F1 evaluation and full program — one line has had its right-hand side removed. Put it back.
def best_valid_span(s_lg, e_lg):
"""Return (start, end) maximizing s[i]+e[j] subject to i<=j."""
T = s_lg.shape[0]
best, bi, bj = float('-inf'), 0, 0
for i in range(T):
for j in range(i, T):
v = s_lg[i].item() + e_lg[j].item()
if v > best:
best, bi, bj = v, i, j
return bi, bj
def token_f1(pred_ids, gold_ids):
ps, gs = set(pred_ids), set(gold_ids)
common = ps & gs
if not common: return 0.0
p = len(common)/len(ps); r = len(common)/len(gs)
return 2pr/(p+r)
model.eval()
total_em, total_f1, n = 0, 0.0, 0
for seed in range(20): # 20 held-out examples
ids, segs, gs, ge = make_batch(B=1, seed=seed+1000)
with torch.no_grad():
s_lg, e_lg = model(ids, segs)
ps, pe = best_valid_span(s_lg[0], e_lg[0])
pred_ids = ids[0, ps:pe+1].tolist()
gold_ids = ids[0, gs[0]:ge[0]+1].tolist()
em = int(pred_ids == gold_ids)
f1 = token_f1(pred_ids, gold_ids)
total_em += em; total_f1 += f1; n += 1
print(f'EM=___ F1=___ (n=___)')
Why: pred_ids is what everything below it consumes, so the wrong expression here fails later and somewhere else. In a real system you would map token indices to character offsets via the tokenizer's offset_mapping.
Worked example
Extract predicted spans using the constrained decoding procedure, then map token indices back to answer strings and compute EM and F1
Why: In a real system you would map token indices to character offsets via the tokenizer's offset_mapping. Here we use synthetic token IDs directly to verify EM/F1 logic.
def best_valid_span(s_lg, e_lg):
"""Return (start, end) maximizing s[i]+e[j] subject to i<=j."""
T = s_lg.shape[0]
best, bi, bj = float('-inf'), 0, 0
for i in range(T):
for j in range(i, T):
v = s_lg[i].item() + e_lg[j].item()
if v > best:
best, bi, bj = v, i, j
return bi, bj
def token_f1(pred_ids, gold_ids):
ps, gs = set(pred_ids), set(gold_ids)
common = ps & gs
if not common: return 0.0
p = len(common)/len(ps); r = len(common)/len(gs)
return 2*p*r/(p+r)
model.eval()
total_em, total_f1, n = 0, 0.0, 0
for seed in range(20): # 20 held-out examples
ids, segs, gs, ge = make_batch(B=1, seed=seed+1000)
with torch.no_grad():
s_lg, e_lg = model(ids, segs)
ps, pe = best_valid_span(s_lg[0], e_lg[0])
pred_ids = ids[0, ps:pe+1].tolist()
gold_ids = ids[0, gs[0]:ge[0]+1].tolist()
em = int(pred_ids == gold_ids)
f1 = token_f1(pred_ids, gold_ids)
total_em += em; total_f1 += f1; n += 1
print(f'EM={total_em/n:.3f} F1={total_f1/n:.3f} (n={n})')Show it off: print the predicted and gold token spans for the first 3 held-out examples
Why: Qualitative inspection confirms the model is pointing at approximately the right region, even if it misses by 1 token in edge cases (the most common error mode for span models).
| step | what to check | success criterion |
|---|---|---|
| model definition | sum(p.numel() for p in model.parameters()) | 101,378 |
| forward pass | s_lg.shape, e_lg.shape | torch.Size([B, T]) each |
| training loss | loss at step 0 | ≈3.0 (= log(T) = log(20)) |
| constrained span | best_valid_span always returns i ≤ j | no inverted spans |
| EM/F1 | EM and F1 > 0 on held-out seeds | model learned something (>random) |
Comparison
Comparison matrix
From Your turn — Milestone 3: EM/F1 evaluation and full program: refill the what to check column from what you know. The rest of the table is as it appeared.
| step | what to check | success criterion |
|---|---|---|
| model definition | sum(p.numel() for p in model.parameters()) | 101,378 |
| forward pass | s_lg.shape, e_lg.shape | torch.Size([B, T]) each |
| training loss | loss at step 0 | ≈3.0 (= log(T) = log(20)) |
| constrained span | best_valid_span always returns i ≤ j | no inverted spans |
| EM/F1 | EM and F1 > 0 on held-out seeds | model learned something (>random) |
Connect it up
Draw it
One page, no notation unless you need it: draw how these connect — SQuAD — the span-prediction task · Two heads on BERT — start & end prediction · EM and F1 — evaluating span predictions · Open-domain QA — retrieve then read. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.
Recap
[CLS] question [SEP] passage [SEP]; output is two token indices — start and end of the answer substring.Linear(d,1) each give (B, T) logits; total new parameters for d=768: 1,538 (0.001% of BERT-base).(CE(start_logits, gold_start) + CE(end_logits, gold_end)) / 2 — cross-entropy over T token-position classes.i ≤ j, maximize s_i + e_j; never take independent argmax blindly.Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.