Exercise Solutions: Masking in LLM Training: Loss, Causal and Padding

Exercise solutions: Masking in LLM Training: Loss, Causal and Padding

These are the worked solutions for the exercises in Part 7, Masking in LLM Training: Loss, Causal and Padding. Read the exercise first; coming here before you have tried it defeats the point.

Exercise 1

The full sequence before padding is the prompt, the response, then the closing token: 11, 12, 13, 21, 22, 23, 99, seven tokens. Padded to length 8 with the pad id:

input_ids      = [11, 12, 13, 21, 22, 23, 99, 0]
attention_mask = [ 1,  1,  1,  1,  1,  1,  1, 0]
labels         = [-100, -100, -100, 21, 22, 23, 99, -100]

Positions 0 to 2, the prompt, are masked to -100 because they are not the assistant’s words. The pad at position 7 is excluded twice over: an attention_mask of zero keeps it out of every attention computation, and a label of -100 keeps it out of the loss. Positions 3 to 6, the response tokens and the closing token, are the only graded span, four tokens out of eight. Decoding that span should read as exactly the assistant’s reply plus nothing else, which is the check this part asks you to run on every real batch.

Exercise 2

Run this part’s sanity check on a training batch, decode the tokens where the label is not -100, and look specifically at whether the last token of that decoded span is the assistant turn’s closing token. A pass looks like the decoded text ending exactly on the closing token, meaning the model was given gradient signal to predict stop at that exact position throughout training. A failure looks like the decoded span ending one token short of the closing token, meaning whatever code built the boundary excluded it, commonly because a slice used something like the last real token minus one, or because the index that located the closing token was off by one. If the closing token was excluded from the graded region, the model was never given a reason to want to emit it, so at inference nothing tells it to stop, which is exactly the rambling behaviour described. The fix lives entirely in how the mask was built, not in anything at inference time.

Exercise 3

Total assistant tokens across the two turns: 12 + 18 = 30. Grading only the final turn keeps the 18 tokens of the second reply and discards all 12 tokens of the first, so the fraction lost is 12 / 30 = 0.4, forty percent of the assistant supervision in this one example. The fraction kept is 18 / 30 = 0.6, sixty percent. Even though this particular example loses forty percent rather than the roughly half this part’s general warning suggests, the underlying point still holds: grading only the last turn always throws away every earlier assistant turn’s tokens in full, and the exact fraction lost simply tracks how the response lengths happen to fall in that example, which is not something you control turn by turn. The safe default is to grade every assistant turn, not to hope the loss happens to be small.

Exercise 4

The causal mask’s rule says nothing about which training example a position belongs to, only that position i may attend to positions at or before i. When examples are concatenated with no boundary reset, position indices simply keep counting upward across the whole packed sequence. If the first example occupies positions 0 to 129 and the second begins at position 130, then a token at position 135, inside the second example, is still at a position after 129, so the causal rule’s own condition, at or before, is satisfied for every position in the first example’s range. From position 135’s point of view the first example is simply more of its own past, and the causal mask has no separate notion of example boundaries to tell it otherwise. This part’s fix is boundary resets: attention and position indices restart at each packed example, so the second example’s positions stop registering the first example’s tokens as being in its past at all.

Exercise 5

The assertion only checks that the batch, taken as a whole, contains at least one label that is not -100 somewhere in it. A single fully truncated example, every one of its own labels at -100, can sit inside an otherwise healthy batch next to several normal examples, and the assertion still passes because those other examples supply plenty of non -100 labels. The truncated example itself still teaches nothing and its forward pass compute is wasted, invisibly. The dataset-level fix removes such examples before batches are ever built, checking each example’s own labels in isolation:

ds = ds.filter(lambda r: any(l != -100 for l in r["labels"]))

Because this runs per example rather than per batch, it does not depend on which other examples happen to be shuffled alongside a bad one, which is the only way to guarantee that no fully masked example ever reaches training.