Lesson 107: Extractive QA — SQuAD, BERT, and Open-Domain Retrieval

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

What this lesson covers

The lesson, slide by slide

1. Extractive QA & SQuAD · BERT · DPR

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.

2. By the end of this lesson you can

Objectives

  1. Describe the SQuAD task and explain why it is a span-prediction problem rather than a generation problem
  2. Implement two linear heads on BERT output (start position, end position) and trace the tensor shapes through the forward pass
  3. Explain why predicting start and end as two independent softmax outputs is a valid approximation and where it breaks down
  4. Compute cross-entropy loss on token-position targets and describe the training signal it provides
  5. Apply the constrained decoding procedure (maximize logit sum over valid start ≤ end pairs) and state when it differs from independent argmax
  6. Compute Exact Match (EM) and token-level F1 for predicted answer spans and interpret their tradeoffs
  7. Describe the Open-Domain QA pipeline: DPR retrieval (dual encoder, dot-product score) → reader (BERT-style span extraction)

3. SQuAD — the span-prediction task

Section

Part 1 of 4

4. SQuAD: from reading comprehension to token indices

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.

fieldexamplenote
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.1always answerable100,837 QA pairs, 536 Wikipedia articles
SQuAD 2.0~33% unanswerableadds 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.

5. Fill in: example for SQuAD: from reading comprehension to token…

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.

fieldexamplenote
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.1always answerable100,837 QA pairs, 536 Wikipedia articles
SQuAD 2.0~33% unanswerableadds adversarial unanswerable questions

6. BERT input for SQuAD: the concatenated sequence

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 rangesegment IDmodel sees
[CLS], q1…qm, [SEP]0question context
p1…pn, [SEP]1passage context
start/end labels—must lie inside passage tokens only

7. What each one costs: BERT input for SQuAD: the concatenated sequence

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 rangesegment IDmodel sees
[CLS], q1…qm, [SEP]0question context
p1…pn, [SEP]1passage context
start/end labels—must lie inside passage tokens only

8. Two heads on BERT — start & end prediction

Section

Part 2 of 4

9. The two-head architecture

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.

10. Break it if you can: The two-head architecture

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.

11. Guess the shape of the answer: Forward pass: shapes and loss — verified

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).

12. Forward pass: shapes and loss — verified

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).

tensorshapeoperation 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_startscalar 3.2292cross_entropy over T=16 classes, gold=[3,7]
loss_endscalar 2.8630cross_entropy over T=16 classes, gold=[5,11]
total lossscalar 3.0461(loss_start + loss_end) / 2

13. Work backwards from the answer: Forward pass: shapes and loss — verified

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.

14. Why two independent softmax outputs work

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.

15. Something is wrong here: taking independent argmax without the start ≤ end…

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.

16. Trap: taking independent argmax without the start ≤ end constraint

Trap

The 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.

The fix

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.

17. EM and F1 — evaluating span predictions

Section

Part 3 of 4

18. Exact Match (EM) and token-level F1

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.

19. Guess the shape of the answer: EM and F1 — four concrete examples

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).

20. EM and F1 — four concrete examples

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.

predictedgoldEMPRF1
albert einsteinalbert einstein11.0001.0001.000
albert einsteineinstein00.5001.0000.667
the united statesunited states of america00.6670.5000.571
424211.0001.0001.000

21. Work backwards from the answer: EM and F1 — four concrete examples

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.

22. Something is wrong here: using character-level overlap instead of token-level F1

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.

23. Trap: using character-level overlap instead of token-level F1

Trap

The 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.

The fix

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.

24. Break it on purpose: using character-level overlap instead of…

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.

25. Open-domain QA — retrieve then read

Section

Part 4 of 4

26. Open-domain QA: the retriever–reader pipeline

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.

  1. Retriever — ranks passages by relevance to the question. Classical: TF-IDF (DrQA, 2017). Neural: DPR dual encoder (2020).
  2. Reader — BERT-style extractive model applied to each top-K retrieved passage; best span across passages is returned.
  3. Top-K — typically K=100 retrieved passages; reader scores spans from all 100, picks the highest logit sum.

27. DPR: dense passage retrieval via dual encoder

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).

28. Guess the shape of the answer: DPR retrieval — dot-product ranking verified

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.

29. DPR retrieval — dot-product ranking verified

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.

rankpassagedot-product score
130.2467
220.2267
340.0991
41-0.0375
50-0.2886

30. Fill in: passage for DPR retrieval — dot-product ranking verified

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.

rankpassagedot-product score
130.2467
220.2267
340.0991
41-0.0375
50-0.2886

31. Generative QA with T5

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.

propertyextractive (BERT)generative (T5)
outputstart/end token indicesauto-regressive token sequence
answermust be a passage substringany string; can be rephrased
unanswerableSQuAD 2.0: add 'no answer' classmodel outputs 'unanswerable'
latencyO(1) decode stepO(answer_len) decode steps
EM/F1standard SQuAD metricsROUGE 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.

32. Break it if you can: Generative QA with T5

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.

33. Without one step: The QA system playbook

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:

  1. Formulate: is the answer a passage substring? → extractive (BERT). Is it generated/rephrased? → generative (T5). Is the passage unknown? → open-domain (DPR…
  2. Tokenize: pack [CLS] question [SEP] passage [SEP] into ≤512 BPE tokens; record passage token offsets for span-to-character mapping.
  3. Encode: run BERT; collect hidden states h ∈ ℝ^{T×d} for the full sequence.
  4. Score: start_logits = Linear_s(h).squeeze(-1), end_logits = Linear_e(h).squeeze(-1) — both shape (T,).
  5. Train: loss = (cross_entropy(start_logits, gold_start) + cross_entropy(end_logits, gold_end)) / 2.
  6. Decode: enumerate all i ≤ j, pick argmax_i,j [s_i + e_j]; mask question tokens and [SEP]/[CLS] tokens out.
  7. Evaluate: normalize both pred and gold (lower, strip punct/articles, split), compute EM (exact equality) and token-level F1 (bag-of-word overlap).
  8. Open-domain: pre-compute passage embeddings offline (DPR E_P), retrieve top-K by dot product, run reader on each, return highest logit-sum span.

34. The QA system playbook

Pattern

  1. Formulate: is the answer a passage substring? → extractive (BERT). Is it generated/rephrased? → generative (T5). Is the passage unknown? → open-domain (DPR + reader).
  2. Tokenize: pack [CLS] question [SEP] passage [SEP] into ≤512 BPE tokens; record passage token offsets for span-to-character mapping.
  3. Encode: run BERT; collect hidden states h ∈ ℝ^{T×d} for the full sequence.
  4. Score: start_logits = Linear_s(h).squeeze(-1), end_logits = Linear_e(h).squeeze(-1) — both shape (T,).
  5. Train: loss = (cross_entropy(start_logits, gold_start) + cross_entropy(end_logits, gold_end)) / 2.
  6. Decode: enumerate all i ≤ j, pick argmax_i,j [s_i + e_j]; mask question tokens and [SEP]/[CLS] tokens out.
  7. Evaluate: normalize both pred and gold (lower, strip punct/articles, split), compute EM (exact equality) and token-level F1 (bag-of-word overlap).
  8. Open-domain: pre-compute passage embeddings offline (DPR E_P), retrieve top-K by dot product, run reader on each, return highest logit-sum span.

35. Where does it stop working: The QA system playbook

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:

  1. Formulate: is the answer a passage substring? → extractive (BERT). Is it generated/rephrased? → generative (T5). Is the passage unknown? → open-domain (DPR…
  2. Tokenize: pack [CLS] question [SEP] passage [SEP] into ≤512 BPE tokens; record passage token offsets for span-to-character mapping.
  3. Encode: run BERT; collect hidden states h ∈ ℝ^{T×d} for the full sequence.
  4. Score: start_logits = Linear_s(h).squeeze(-1), end_logits = Linear_e(h).squeeze(-1) — both shape (T,).
  5. Train: loss = (cross_entropy(start_logits, gold_start) + cross_entropy(end_logits, gold_end)) / 2.
  6. Decode: enumerate all i ≤ j, pick argmax_i,j [s_i + e_j]; mask question tokens and [SEP]/[CLS] tokens out.
  7. Evaluate: normalize both pred and gold (lower, strip punct/articles, split), compute EM (exact equality) and token-level F1 (bag-of-word overlap).
  8. Open-domain: pre-compute passage embeddings offline (DPR E_P), retrieve top-K by dot product, run reader on each, return highest logit-sum span.

36. Rule out three: Check 1 — head architecture

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.

  • A. 1,538 (768×2 weights + 2 biases)
  • B. 1,536 (768×2 weights only, no biases)
  • C. 768 (one shared linear layer)
  • D. 589,824 (768×768 per head)

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%.

37. Check 1 — head architecture

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?

  • A. 1,538 (768×2 weights + 2 biases) (correct)
  • B. 1,536 (768×2 weights only, no biases)
  • C. 768 (one shared linear layer)
  • D. 589,824 (768×768 per head)

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%.

Why B tempts people
Forgets the bias term; PyTorch nn.Linear has bias=True by default, adding 1 scalar per output unit.
Why C tempts people
BERT-SQuAD uses two separate heads — one predicts start, one predicts end — not a single shared layer. Sharing would conflate the two tasks.
Why D tempts people
Confuses a square weight matrix (768×768) with the correct (768×1) output projection; the QA head projects to a single logit, not to d dimensions.

38. Answer it before you see the options: Check 2 — EM vs F1

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.

39. Check 2 — EM vs F1

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)?

  • A. EM = 0, F1 = 0.800 (correct)
  • B. EM = 0, F1 = 0.667
  • C. EM = 1, F1 = 1.000
  • D. EM = 0, F1 = 1.000

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.

Why B tempts people
Computes P = 2/3 correctly but uses F1 = 2·P·R/(P+R) = 2·(2/3)·(2/3)/((2/3)+(2/3)) = 0.667, incorrectly taking R = 2/3 instead of R = 2/2 = 1. Recall is common/gold, not common/pred.
Why C tempts people
Incorrectly concludes EM=1 because 'new york' is a substring of 'new york city'; EM requires exact string equality after normalization, not substring match.
Why D tempts people
Sets F1=1 because recall R=1 (both gold tokens are in the prediction), but ignores that precision < 1 since 'city' is a spurious extra token; harmonic mean of P<1 and R=1 is always <1.

40. Rule out three: Check 3 — training loss interpretation

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.

  • A. It teaches the model to maximize P(start=i) and P(end=j) independently, which is valid because both heads share BERT's context vectors that jointly encode passage semantics
  • B. It teaches the model to predict the joint span distribution P(start=i, end=j) exactly, which requires a T×T output matrix
  • C. It teaches start position only; the end head is supervised by the start head's output, making them dependent
  • D. It maximizes the log-probability of the correct character offsets, not token indices, so the loss is computed in character space

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.

41. Check 3 — training loss interpretation

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?

  • A. It teaches the model to maximize P(start=i) and P(end=j) independently, which is valid because both heads share BERT's context vectors that jointly encode passage semantics (correct)
  • B. It teaches the model to predict the joint span distribution P(start=i, end=j) exactly, which requires a T×T output matrix
  • C. It teaches start position only; the end head is supervised by the start head's output, making them dependent
  • D. It maximizes the log-probability of the correct character offsets, not token indices, so the loss is computed in character space

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.

Why B tempts people
A joint distribution over all (start, end) pairs would require a T² output and a joint softmax — far more parameters and intractable for T=512. The factored formulation with two T-way softmaxes is the efficiency trick that makes BERT-SQuAD practical.
Why C tempts people
The end head is supervised directly by the gold end label, not by the start head's output. The two heads are trained in parallel with independent losses; there is no sequential dependency at training time.
Why D tempts people
SQuAD labels are stored as character offsets in the dataset, but they are converted to token indices before training. The cross-entropy loss operates over the discrete token vocabulary of positions 0…T-1, not in character space.

42. Your turn — build a QA system

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.

43. Your turn — Milestone 1: model definition and forward pass

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,378

Forward 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.

componentoutput shapeparam 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

44. Which is which, by output shape

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.

(B, T, 64)
tok_embed Embedding(500,64); pos_embed Embedding(32,64); seg_embed Embedding(2,64); MHSA (per Block): qkv+out; FFN (per Block): 64→128→64
—
LayerNorm (2 per block + 1 final); TOTAL
(B, T, 2)
QA head Linear(64,2)
g1
output shape is "(B, T, 64)" for tok_embed Embedding(500,64), pos_embed Embedding(32,64), seg_embed Embedding(2,64), MHSA (per Block): qkv+out, FFN (per Block): 64→128→64 — that is what the table on "Your turn — Milestone 1: model…" records, and it is the single property separating this group from the rest.
g2
output shape is "—" for LayerNorm (2 per block + 1 final), TOTAL — that is what the table on "Your turn — Milestone 1: model…" records, and it is the single property separating this group from the rest.
g3
output shape is "(B, T, 2)" for QA head Linear(64,2) — that is what the table on "Your turn — Milestone 1: model…" records, and it is the single property separating this group from the rest.

45. Predict the next row: Your turn — Milestone 2: training loop with…

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

stepexpected lossinterpretation
0~3.0uniform over 20 classes: log(20) ≈ 3.0
50~2.0model narrowing to correct region
100~1.0substantially learned start/end offsets
150~0.5near-convergence on synthetic task
200~0.2overfitting 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.

46. Your turn — Milestone 2: training loop with synthetic data

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.

stepexpected lossinterpretation
0~3.0uniform over 20 classes: log(20) ≈ 3.0
50~2.0model narrowing to correct region
100~1.0substantially learned start/end offsets
150~0.5near-convergence on synthetic task
200~0.2overfitting to synthetic random seeds

47. What each one costs: Your turn — Milestone 2: training loop with…

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.

stepexpected lossinterpretation
0~3.0uniform over 20 classes: log(20) ≈ 3.0
50~2.0model narrowing to correct region
100~1.0substantially learned start/end offsets
150~0.5near-convergence on synthetic task
200~0.2overfitting to synthetic random seeds

48. Restore the missing line: Your turn — Milestone 3: EM/F1 evaluation and…

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.

49. Your turn — Milestone 3: EM/F1 evaluation and full program

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).

stepwhat to checksuccess criterion
model definitionsum(p.numel() for p in model.parameters())101,378
forward passs_lg.shape, e_lg.shapetorch.Size([B, T]) each
training lossloss at step 0≈3.0 (= log(T) = log(20))
constrained spanbest_valid_span always returns i ≤ jno inverted spans
EM/F1EM and F1 > 0 on held-out seedsmodel learned something (>random)

50. Fill in: what to check for Your turn — Milestone 3: EM/F1 evaluation…

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.

stepwhat to checksuccess criterion
model definitionsum(p.numel() for p in model.parameters())101,378
forward passs_lg.shape, e_lg.shapetorch.Size([B, T]) each
training lossloss at step 0≈3.0 (= log(T) = log(20))
constrained spanbest_valid_span always returns i ≤ jno inverted spans
EM/F1EM and F1 > 0 on held-out seedsmodel learned something (>random)

51. Connect it up: Lesson 107: Extractive QA — SQuAD, BERT, and Open-Domain Retrieval

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.

52. Lesson 107 — key takeaways

Recap

Sources

  1. USAAIO Year-Long Master Lesson Plan, Lesson 107 — Extractive QA, SQuAD, BERT — Barron · USAAIO Round 2 Preparation, 2026
  2. Rajpurkar et al. 'SQuAD: 100,000+ Questions for Machine Comprehension of Text' (EMNLP 2016) — arXiv:1606.05250
  3. Karpukhin et al. 'Dense Passage Retrieval for Open-Domain Question Answering' (EMNLP 2020) — arXiv:2004.04906
  4. TinyBERTforQA from scratch, start/end head shapes, loss values, DPR similarity, EM/F1 on toy strings, span constraint — all verified with torch 2.7.1+cpu and numpy 2.2.6, June 2026 — Real execution, verified

Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.

Book on Wyzant · Text (657) 465-8108