Lesson 4 of 4 · 35 min

Prove the optimized implementation matches a reference

Design a small reference test for a vectorized operation.

Research code often replaces a clear implementation with a faster one. The replacement must preserve the intended computation before its speed matters. Keep a simple reference that is easy to inspect. Then compare the optimized result on ordinary inputs and edge cases designed to expose indexing, masking, and reduction mistakes.
Choose tolerances based on the computation and numeric type. Floating-point operations are not associative, so changing reduction order can produce small differences. A tolerance should allow expected rounding without hiding a semantic error. Use both absolute and relative tolerance where appropriate, and inspect worst-case differences. A single mean error can conceal one badly wrong element.
Shape checks are necessary but insufficient. Two arrays can have the same shape while one reduces over the wrong axis. Masks can broadcast without an exception and still include padded tokens. Use asymmetric small data where wrong axes produce visibly different results. Symmetric all-ones fixtures are often too forgiving.
The reference must be independent enough to detect the defect. Copying the same indexing expression into both implementations can reproduce the same error. For a complex kernel, compare against a trusted higher-level operation and separately check mathematical invariants. Keep the evaluation tests outside the optimization target so speed gains cannot come from weakening correctness.

Worked example

A teaching masked mean has values [2, 4, 100] and mask [1, 1, 0]. The intended output is (2 + 4) / 2 = 3. A buggy implementation zeros the padded value but divides by total length 3, producing 2. Another ignores the mask and produces about 35.33.
A two-row fixture adds [1, 9, 20] with mask [1, 0, 0], whose output is 1. These rows expose denominator and axis errors. The contract also defines an all-masked row: return a documented neutral value with an invalid-row flag, or reject it. Dividing by zero and letting NaN propagate silently is not a complete design.

Exercise and solution

For values [3, 6, 9, 12] and mask [1, 0, 1, 0], calculate the correct mean and two likely wrong outputs.
The correct result is 6. Dividing the masked sum by four gives 3; ignoring the mask gives 7.5. Award one point each for these results, an explicit all-masked policy, and a test that compares the optimized path to an independently clear reference. Explain why using [1, 1, 1, 1] alone would make several bugs harder to see.

Lab artifact: an independent reference

The reference below is intentionally simple. It validates the row mask and handles an all-masked row explicitly. The neutral value is paired with a validity flag, so downstream aggregation can exclude invalid rows according to a declared policy rather than treating the zero as an observed loss.
python
1def masked_row_mean(values, mask):2    assert len(values) == len(mask)3    assert all(m in (0, 1) for m in mask)4    valid = [v for v, m in zip(values, mask) if m == 1]5    if not valid:6        return 0.0, False7    return sum(valid) / len(valid), True
A vectorized implementation should be checked against this reference on a set that includes different valid lengths, negative values if the operation permits them, large excluded padding, and an empty valid set. Use a tolerance appropriate to the dtype. The reference here assumes finite numeric values; a production contract must say what happens with NaNs or infinities. Multiplying an excluded NaN by zero can still produce NaN, so “zero masked entries by multiplication” is not equivalent to selecting finite valid entries for every possible input.

A second failure case: row mean versus global token mean

Consider two rows. Row A has one valid loss of ten. Row B has three valid losses of two, two, and two. The mean of valid row means is (10 + 2) / 2 = 6. The global valid-token mean is (10 + 2 + 2 + 2) / 4 = 4. Both are defined objectives. They weight sequences differently.
ReductionNumeratorDenominatorResult
Equal weight per valid row10 + 22 rows6
Equal weight per valid token164 tokens4
If an optimized kernel changes from token weighting to row weighting, a small absolute-error tolerance should not excuse the difference. It is a semantic change. The intended loss contract must name the unit of weighting. PyTorch's loss documentation describes reduction and target conventions, but the exact masked multirow objective in this workbook is an original fixture.
The same distinction affects gradients. A row with one valid token receives greater weight relative to each token in a long row under equal-row averaging. A test that checks only the scalar on equal-length rows may miss the change. Unequal valid lengths are essential when the bug concerns the denominator.

Exercise: design a minimal discriminating suite

An optimized function returns the right value on [1, 1, 1] with mask [1, 1, 1]. Choose three additional fixtures and the failures they expose.
Use [2, 4, 100] with [1, 1, 0] to expose ignored masks and padded denominators. Use unequal valid lengths across two rows, as above, to expose row-versus-token weighting. Use a completely masked row to test the declared invalid-row policy. A fourth useful case uses a non-finite excluded value if the input contract allows it, checking whether masking semantics actually exclude it. Award one point for each fixture, one for the mapped defect, and one for choosing a declared objective rather than assuming one.

Misconceptions to correct

“Matching shape and finite loss establish equivalence” fails because axis and weighting bugs often produce plausible finite scalars. “A larger tolerance solves low-precision mismatch” fails when the discrepancy is systematic and much larger than expected rounding. Examine the difference pattern, compare a higher-precision reference where useful, and justify the tolerance from numerical behavior.
After correctness, benchmark the same workload and resource contract. The historical Anthropic performance task is useful evidence that some employer tasks explicitly protect correctness tests while optimizing execution. It is not a claim that this masked-mean fixture appeared in that task or that today's hiring process uses the same benchmark. Keep employer provenance separate from the original lesson.
The final report should contain the input fixture, intended reduction, expected value, observed values, tolerance, and result for each case. Another engineer should be able to tell whether a speed improvement preserved the computation or quietly changed the training objective.

Interview probe

Original practice: A vectorized loss is faster. How do you establish it is the same loss? A strong answer uses a reference, asymmetric fixtures, mask and reduction cases, and justified tolerances. Follow up with an all-masked batch. A weak answer checks only output shape and average training loss.

Sources

docsPyTorch: automatic mixed precision examplesdocs.pytorch.orgdocsNeurIPS: paper checklistneurips.ccdocsAnthropic: historical performance take-homegithub.comdocsPyTorch: cross-entropy loss and reductiondocs.pytorch.org

Checkpoint

A masked mean divides the valid-value sum by padded length. What changed?

AOnly the numeric rounding error; the objective is preserved.BOnly a constant scale independent of valid sequence length.CThe mathematical denominator and therefore the objective.DNothing if the scalar stays finite.
Sign up free to answer and see why

Checkpoint

One row has valid loss [10], another [2,2,2]. Global valid-token mean equals?

A6B4C8D2
Sign up free to answer and see why

Checkpoint

An excluded value is NaN. Why can multiplication by a zero mask fail to exclude it numerically?

ANaN times zero can remain NaN.BConverting the binary mask to the value dtype removes the NaN.CDividing by the valid count removes non-finite padding contributions.DUsing a wider floating-point dtype makes zero times NaN finite.
Sign up free to answer and see why

Checkpoint

Which suite best detects wrong-axis and mask-reduction defects?

AOne all-ones tensor with no padding.BShape checks and average loss over a large run.COnly fixtures where every row has the same valid length.DAsymmetric values, unequal valid lengths, and explicit all-masked behavior.
Sign up free to answer and see why

Checkpoint

An optimized loss differs systematically by row length. What should happen before increasing tolerance?

AAccept it if training converges.BCheck whether row versus token weighting changed.CAverage away the discrepancy over many steps.DMeasure speed first and validate only the fastest path.
Sign up free to answer and see why

Can you make a minimal fixture expose a mask, axis, or weighting defect? Rate confidence from 1 to 5 and explain how you chose the tolerance and invalid-row policy.

Not yetGetting thereConfident

Wrap-up

  • Keep a clear reference and adversarially small fixtures. Correctness is a prerequisite to meaningful optimization.

Sources

Free to read · better with Enzo

Learn it with Enzo

Save your progress, answer the checkpoints, and let Enzo quiz you on what you just read.