Reading and Reimplementing Papers

Turning an equation into tensor code

Transcribe the paper line by line with the paper's own variable names, then prove the transcription right against a trusted implementation before you build anything on top of it.

Read these first

On this page 7
  1. Why the line-by-line rule
  2. The rule that saves the most time
  3. How it works
  4. Prove it, do not eyeball it
  5. A real example you have seen
  6. Remember this
  7. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

Turning an equation into code means writing one line of code per line of the paper, keeping the paper's own names, and then proving the result matches something you trust.

Picture assembling a flat-pack cupboard. The instructions show ten steps. You do not read all ten, form a general impression, and start screwing things together. You do step one, look at the picture, check the panel is facing the right way, and only then do step two.

Skipping ahead is how people end up with a cupboard whose door opens into the wall.

Why the line-by-line rule

The temptation is to read the whole method, understand it, and then write your own clean version. That feels more intelligent. It goes wrong for one reason: when the numbers come out different, you have no idea where.

A line-by-line transcription is ugly, has variable names like m and v, and looks nothing like the code you would write for a colleague. It has one enormous advantage. Every line points at a line in the paper, so a disagreement can be traced to one place.

Make it correct first. Make it beautiful second, if ever.

The rule that saves the most time

Keep the paper's names.

If the paper calls something theta, call it theta. Not params, not weights, not w. The moment you rename things, you are holding two vocabularies in your head at once and translating between them on every line. That is where mistakes enter.

You can rename everything later, once the code is proven right. Renaming proven code is safe. Renaming unproven code is how a bug becomes permanent.

How it works

   paper line 1  →  code line 1  ─┐
   paper line 2  →  code line 2   │
   paper line 3  →  code line 3   ├→  run it  →  compare against a
   paper line 4  →  code line 4   │              trusted implementation
   paper line 5  →  code line 5  ─┘
                                        agree to many decimals? done.
                                        disagree? one line is wrong,
                                        and you can find which.

Prove it, do not eyeball it

The last box is the part people skip, and it is the part that matters.

"The loss went down, so it works" is not evidence. Many wrong implementations make the loss go down. Wrong step sizes, wrong signs on small terms and missing corrections all still train, a bit worse, silently.

The real check is a number. Run your version and a trusted version on the same input, and compare the outputs digit by digit. Agreement to ten or more decimal places means they are the same computation. Anything less means they are not.

A real example you have seen

Every calculator app was tested this way. Nobody shipped one because the answers looked about right. They ran millions of inputs against a reference and compared. Your transcription of a paper deserves the same standard, and it takes about five minutes to set up.

Remember this

  • One line of code per line of the paper, in the paper's own names.
  • Ugly and traceable beats elegant and unverifiable.
  • Compare against a trusted implementation on the same input. Numbers, not impressions.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Verified with torch 2.5.1 and numpy 1.26.4 on CPU. Runs instantly.

Algorithm 1, transcribed

Kingma and Ba's Adam is ten lines of pseudocode. Here it is as ten lines of numpy, with the paper's names kept exactly, and then checked against PyTorch's implementation on the same problem.

adam_transcribed.py
import numpy as np
import torch

def adam_from_paper(grad_fn, theta, alpha=0.001, beta1=0.9, beta2=0.999,
                    eps=1e-8, steps=100):
    """Algorithm 1 of Kingma and Ba (2015), line for line."""
    m = np.zeros_like(theta)
    v = np.zeros_like(theta)
    for t in range(1, steps + 1):                  # t starts at 1, not 0
        g = grad_fn(theta)
        m = beta1 * m + (1 - beta1) * g            # first moment
        v = beta2 * v + (1 - beta2) * g * g        # second moment, elementwise
        m_hat = m / (1 - beta1 ** t)               # bias correction
        v_hat = v / (1 - beta2 ** t)
        theta = theta - alpha * m_hat / (np.sqrt(v_hat) + eps)
    return theta

scales = np.array([1.0, 100.0, 0.01])
mine = adam_from_paper(lambda th: 2 * scales * th, np.ones(3), alpha=0.01, steps=100)

w = torch.ones(3, dtype=torch.float64, requires_grad=True)
opt = torch.optim.Adam([w], lr=0.01)
s = torch.tensor(scales, dtype=torch.float64)
for _ in range(100):
    opt.zero_grad()
    (s * w**2).sum().backward()
    opt.step()

print("from the paper :", mine)
print("torch.optim.Adam:", w.detach().numpy())
print("largest disagreement:", np.abs(mine - w.detach().numpy()).max())
Output
from the paper : [0.22444605 0.22444604 0.22444637]
torch.optim.Adam: [0.22444605 0.22444604 0.22444637]
largest disagreement: 5.551115123125783e-17

A disagreement of 5.6e-17 after a hundred steps is floating-point dust — the two computations are the same. That single line of output is worth more than any amount of staring at the code.

Note what this test does not need: a dataset, a model, a GPU, or a training run. It needs one function whose gradient you can write down.

The walkthrough

float64 on the PyTorch side is deliberate. numpy defaults to double precision and PyTorch to single. Compare a float32 run against a float64 run and you get disagreements around 1e-7, which are indistinguishable from a real bug in the last decimals. Match the dtypes first, then interpret the difference.

t starts at 1. The paper's loop is one-based, and $t$ appears inside $\beta_1^t$. Start at 0 and the first bias correction divides by $1 - \beta_1^0 = 0$. This is the single most common transcription error in this algorithm.

m and v live outside the loop. They are state carried between steps. Recreating them inside the loop turns Adam into something that looks similar and behaves nothing like it — and the loss will still go down, which is why this bug survives.

g * g is elementwise, not a dot product. Whenever a paper squares a vector, decide whether it means elementwise or an inner product. Here the shapes tell you: v must match theta, so the square is elementwise.

A quadratic is the right test problem. Its gradient is exact and cheap, the minimum is known, and it exercises the per-coordinate behaviour with three different scales. Testing an optimiser against a real network first would mean debugging two things at once.

When there is no reference to compare against

Most methods you reimplement will not exist in a library. Three substitutes, in order of strength.

Compare a fast version against a slow one. Write the plain nested-loop version first, then the vectorised version, and assert they agree on small inputs. The loop version is too slow to use and too simple to be wrong.

Check gradients numerically. If the paper defines a loss and its derivative, verify the derivative against finite differences — see verifying gradients with gradcheck.

Check the properties the maths promises. Probabilities sum to one. A symmetric kernel produces a symmetric matrix. A normalised embedding has length one. An attention row sums to one. Each of these is one assertion and catches a whole class of errors.

Common mistakes

Rewriting the algorithm "more cleanly" during transcription. Fusing two lines, reordering updates, dropping a term that looks redundant. Every simplification is an unverified change. Transcribe, verify, then refactor.

Comparing against a library with different defaults. PyTorch's Adam defaults to eps=1e-8 outside the square root, matching Algorithm 1. Other frameworks place epsilon differently, and some apply weight decay by a different route. Read the reference implementation's own formula before calling a mismatch a bug in your code.

Testing only on the happy path. Try a zero gradient, a huge gradient, and a single parameter. Division-by-zero and broadcasting bugs surface at the edges, not in the middle.

Declaring success from a loss curve. A curve going down is compatible with a dozen wrong implementations. Insist on a numerical comparison somewhere in the chain, even if it is against your own slow version.

Try it yourself

Break the transcription on purpose in four ways: start t at 0 with the divide-by-zero guarded, delete both bias corrections, move m = np.zeros_like(theta) inside the loop, and change g * g to g @ g. Predict which of the four still runs, which still converges, and which produces a disagreement small enough to miss.

What to learn next

Researcher — Mathematics and papers.

Equivalence testing as the unit of trust

A reimplementation is a claim of equivalence between two computations. Treat it as a claim to be tested at a stated tolerance, not as a state of belief.

Practical tolerances: agreement at atol=1e-12 in float64 indicates identical arithmetic up to summation order; around 1e-6 to 1e-7 in float32 indicates the same computation at single precision; anything above 1e-4 indicates a real difference in the algorithm, not in the arithmetic. Reporting the tolerance you tested at, and the dtype, makes the claim reproducible.

For stochastic components, equivalence is tested in distribution rather than pointwise: match generator seeds where the reference exposes them, otherwise compare means and variances of the output over many draws with a two-sample test.

Literal transcription can be numerically wrong

Mathematical notation is indifferent to floating point. Code is not. The standard rewrites:

  • Softmax and log-sum-exp. $\log \sum_j e^{z_j}$ overflows for $z_j > 709$ in double precision. Subtract $\max_j z_j$ first; the result is algebraically identical and numerically stable.
  • Cross-entropy from probabilities. Computing $-\log p$ after an explicit softmax loses precision and can produce infinities at $p = 0$. Frameworks combine the two into one stable operation, which is why library losses take logits rather than probabilities. See logits and loss pitfalls.
  • Variance in one pass. $\mathbb{E}[X^2] - (\mathbb{E}[X])^2$ is catastrophic cancellation waiting to happen; Welford's algorithm is the stable form.
  • $\log(1+x)$ and $e^x - 1$ for small $x$ need log1p and expm1.
  • Near-zero denominators. Adam's $\epsilon$ exists for exactly this reason. When you transcribe a division, ask what happens when the denominator is small, and check whether the paper's guard is inside or outside a square root.

The general rule: transcribe literally, verify against a reference, and only then substitute a numerically stable rewrite — re-running the equivalence test after the substitution.

Vectorisation as a refactor, not a rewrite

The safe sequence is loops first, tensors second, and an equivalence test between them at every step. Index-heavy expressions translate most reliably through einsum, where the subscript string is an explicit, checkable statement of which indices are summed. Once the vectorised form matches the loop form on small random inputs, scale up.

Two properties worth measuring rather than assuming after vectorisation: peak memory, since a vectorised form frequently materialises an intermediate of size $O(n^2)$ that the loop never held; and whether the operation is memory-bound or compute-bound, which determines whether further optimisation is worth attempting at all. See timing GPU code correctly.

What to record

The artefact of this stage is not only code. Keep, alongside the implementation: the equation numbers each function corresponds to, the reference used for verification, the tolerance achieved, the dtype, and the list of properties asserted. When your results later disagree with the paper's — the subject of why your reimplementation is three points worse — this record is what lets you rule the core computation out of the investigation in five minutes instead of five days.

What to learn next