Skip to content

Newton's method

ModuleM01.2 · build · Python · Pass 2 · 2 to 3 h
You buildpython/tinyllm/num/newton.py: newton (roots of a scalar function, with a convergence report) and rsqrt_newton (the division-free iteration for 1/x1/\sqrt{x})
Contractcourse/contracts/py/tinyllm/num/newton.pyi
Testscourse/tests/M01.2/test_newton.py (what they check: section 4)
NeedsM01.1 central_diff, the slope Newton uses when you pass no derivative (or --ref-deps)
Used byM07.7’s fit_temperature (temperature scaling) solves for 1/T1/T with newton · optional M09.5 ports rsqrt_newton to C as the reciprocal square root inside L9.6’s RMSNorm kernel and checks it against your Python · optional M10.6 uses Newton-Schulz iterations for Muon (section 6)
MilestoneMS-P2 (the Pass 2 gate)
Optional depthOpenStax, Calculus Volume 1 (free), section 4.9 (Newton’s method); Sauer, Numerical Analysis, sections 1.4 and 1.5 (convergence order, when Newton fails); Lomont, “Fast inverse square root” (2003) for the magic constant
  • Newton’s method replaces ff by its tangent line at xnx_n and jumps to the tangent’s zero: xn+1=xn−f(xn)/f′(xn)x_{n+1} = x_n - f(x_n)/f'(x_n) (test_hand_example).
  • Near a simple root it converges quadratically: en+1≈Cen2e_{n+1} \approx C e_n^2 with C=∣f′′/(2f′)∣C = |f''/(2f')| at the root, so the number of correct digits doubles every step (test_digits_double_per_step).
  • It is only locally convergent: it can cycle, diverge, or hit a flat tangent. A correct implementation reports each of these instead of returning a wrong number (test_cycle_raises_after_max_iter, test_divergence_raises, test_zero_derivative_raises).
  • The stopping test must be relative, ∣xn+1−xn∣≤tol⋅max⁡(1,∣xn+1∣)|x_{n+1} - x_n| \le \mathrm{tol} \cdot \max(1, |x_{n+1}|), because floats near a large root are far apart (test_relative_tolerance_for_large_roots).
  • Applied to g(y)=1/y2−xg(y) = 1/y^2 - x, Newton computes 1/x1/\sqrt{x} with only multiplications, y←y(32−12xy2)y \leftarrow y(\frac{3}{2} - \frac{1}{2} x y^2); two steps from a 3.5 percent guess reach float32 accuracy (test_rsqrt_error_squares_each_step).
Terminal window
ol start M01.2 # stubs python/tinyllm/num/newton.py into your repo
ol tests M01.2 # read the test catalog first: rung R0, you write no tests here
ol check M01.2 # exit code is the verdict
ol check M01.2 --ref-deps # only if your M01.1 is not passing yet
ol diff M01.2 # after passing: your code against the reference

Every transformer you build normalizes its activations, and from L7.1 on the normalization is RMSNorm: divide a vector by mean of squares\sqrt{\text{mean of squares}}. In Pass 6 your C kernel (L9.6) computes it for every token of every layer, and a division plus a square root per element is the slow way. Hardware and libraries compute 1/x1/\sqrt{x} directly, by a cheap guess refined with Newton’s method, and M09.5 does the same in C. This module teaches the method on paper and in Python, where you can see the digits double, so that the C port has a trusted reference: your own Python, step for step. The general newton you write here is also the first algorithm in the course that can fail in ways a test must catch: a flat tangent, a cycle, a runaway. You build it to say so.

SymbolMeaningType / shape
ffa differentiable function of one real variableCallable[[float], float]
f′f'its derivative (df), or central_diff from M01.1 when df is NoneCallable[[float], float]
x⋆x^\stara root: f(x⋆)=0f(x^\star) = 0float
xnx_nthe nn-th iterate, x0x_0 the starting guessfloat
en=xn−x⋆e_n = x_n - x^\starthe error of the nn-th iteratefloat
tolthe relative step tolerancefloat, default 10−1210^{-12}
xx (in 2.4)a positive number whose 1/x1/\sqrt{x} we wantfloat32 or float64 array
yky_kthe kk-th estimate of 1/x1/\sqrt{x}same dtype as xx

A root of ff is a number x⋆x^\star with f(x⋆)=0f(x^\star) = 0. Many quantities are roots in disguise: 2\sqrt{2} is the positive root of x2−2x^2 - 2, ln⁡3\ln 3 is the root of ex−3e^x - 3, and 1/a1/\sqrt{a} is a root of 1/y2−a1/y^2 - a. A root is simple when f′(x⋆)≠0f'(x^\star) \neq 0: the graph crosses zero at an angle instead of touching it.

Near xnx_n the tangent line approximates ff (that is what the derivative means, M01.1):

f(x)≈f(xn)+f′(xn)(x−xn).f(x) \approx f(x_n) + f'(x_n)(x - x_n) .

The tangent line is zero at x=xn−f(xn)/f′(xn)x = x_n - f(x_n)/f'(x_n), and Newton’s method takes that point as the next iterate:

xn+1=xn−f(xn)f′(xn).x_{n+1} = x_n - \frac{f(x_n)}{f'(x_n)} .

For f(x)=x2−2f(x) = x^2 - 2 the step is xn+1=xn−xn2−22xn=12(xn+2xn)x_{n+1} = x_n - \frac{x_n^2 - 2}{2x_n} = \frac{1}{2}\left(x_n + \frac{2}{x_n}\right): average your guess with 2 divided by it, the method the Babylonians used for square roots. If f′(xn)=0f'(x_n) = 0 the tangent is flat and never crosses zero, so there is no next iterate.

How fast does the error shrink? Expand ff around xnx_n and evaluate at the root (Taylor’s formula with the exact remainder, M02.1): for some ξ\xi between xnx_n and x⋆x^\star,

0=f(x⋆)=f(xn)+f′(xn)(x⋆−xn)+12f′′(ξ)(x⋆−xn)2.0 = f(x^\star) = f(x_n) + f'(x_n)(x^\star - x_n) + \tfrac{1}{2} f''(\xi)(x^\star - x_n)^2 .

Divide by f′(xn)f'(x_n) and use the definition of xn+1x_{n+1}: xn+1−x⋆=f′′(ξ)2f′(xn)(xn−x⋆)2x_{n+1} - x^\star = \frac{f''(\xi)}{2 f'(x_n)} (x_n - x^\star)^2, that is

en+1=f′′(ξ)2f′(xn) en2≈C en2,C=∣f′′(x⋆)2f′(x⋆)∣.e_{n+1} = \frac{f''(\xi)}{2 f'(x_n)}\, e_n^2 \approx C\, e_n^2, \qquad C = \left|\frac{f''(x^\star)}{2 f'(x^\star)}\right| .

The new error is proportional to the square of the old one. If ∣en∣=10−3|e_n| = 10^{-3} and C≈1C \approx 1, then ∣en+1∣≈10−6|e_{n+1}| \approx 10^{-6} and ∣en+2∣≈10−12|e_{n+2}| \approx 10^{-12}: the number of correct digits doubles each step. For x2−2x^2 - 2, C=22⋅22=122≈0.354C = \frac{2}{2 \cdot 2\sqrt{2}} = \frac{1}{2\sqrt 2} \approx 0.354, and the ratio en+1/en2e_{n+1}/e_n^2 approaches exactly that (test_digits_double_per_step).

The proof needs f′(xn)≠0f'(x_n) \neq 0 and xnx_n already close to the root. Far from it nothing is promised. f(x)=x3−2x+2f(x) = x^3 - 2x + 2 started at 0 jumps to 1 and back to 0 forever. f(x)=x3f(x) = \sqrt[3]{x} has its tangent at xx cross zero at −2x-2x, so every step doubles the distance. A method that is only locally convergent must be given an iteration budget, max_iter, and must report when it runs out.

The derivative must move with the iterate. Computing f′(x0)f'(x_0) once and reusing it (the “chord method”) still converges, but only linearly: the error shrinks by a constant factor ∣1−f′(x⋆)/f′(x0)∣|1 - f'(x^\star)/f'(x_0)| per step, and if that factor exceeds 1 it diverges. When you pass no derivative, newton calls central_diff at each xnx_n; its relative error near 10−1010^{-10} keeps the convergence effectively quadratic.

Stopping. Stop when the step is negligible relative to the iterate, ∣xn+1−xn∣≤tol⋅max⁡(1,∣xn+1∣)|x_{n+1} - x_n| \le \mathrm{tol} \cdot \max(1, |x_{n+1}|). An absolute test fails for large roots: floats near 12649 are 1.8×10−121.8 \times 10^{-12} apart, so near 1.6×108\sqrt{1.6 \times 10^8} the last steps hop between two neighbouring floats and an absolute 10−1210^{-12} is never met. The max⁡(1,⋅)\max(1, \cdot) keeps the test absolute near zero, where a relative one would demand impossible precision. And if f(xn)f(x_n) is exactly 0, xnx_n is a root and no step is needed.

Apply the step to g(y)=1y2−xg(y) = \frac{1}{y^2} - x, whose positive root is y=1/xy = 1/\sqrt{x}. With g′(y)=−2/y3g'(y) = -2/y^3:

yk+1=yk−1/yk2−x−2/yk3=yk+yk−xyk32=yk(32−12 x yk2).y_{k+1} = y_k - \frac{1/y_k^2 - x}{-2/y_k^3} = y_k + \frac{y_k - x y_k^3}{2} = y_k \left(\frac{3}{2} - \frac{1}{2}\, x\, y_k^2\right) .

No division and no square root: three multiplications and a subtraction. (The more obvious g(y)=y2−1/xg(y) = y^2 - 1/x needs 1/x1/x first, a division.) Write yk=(1+δk)/xy_k = (1 + \delta_k)/\sqrt{x}, a relative error δk\delta_k. Substituting,

δk+1=−32δk2−12δk3,\delta_{k+1} = -\tfrac{3}{2}\delta_k^2 - \tfrac{1}{2}\delta_k^3 ,

so the relative error squares (times 32\frac{3}{2}) each step. A good first guess comes from the bits of a float32: the integer 0x5f3759df−(bits(x)≫1)\mathtt{0x5f3759df} - (\text{bits}(x) \gg 1), read back as a float, is within 3.5 percent of 1/x1/\sqrt{x} for every positive normal xx, because halving the bits roughly halves the exponent. Then δ1≤32(0.035)2≈1.8×10−3\delta_1 \le \frac{3}{2}(0.035)^2 \approx 1.8 \times 10^{-3} and δ2≈5×10−6\delta_2 \approx 5 \times 10^{-6}: float32 accuracy (its spacing is 6×10−86 \times 10^{-8} relative) after two steps. M09.5 uses exactly this, so rsqrt_newton computes in the dtype it is given: float32 in, float32 arithmetic, float32 out.

2\sqrt{2} from x0=1x_0 = 1, with f(x)=x2−2f(x) = x^2 - 2 and the step xn+1=12(xn+2/xn)x_{n+1} = \frac{1}{2}(x_n + 2/x_n):

nnxnx_nas a fractionen=xn−2e_n = x_n - \sqrt 2en/en−12e_n / e_{n-1}^2
0111−4.1×10−1-4.1 \times 10^{-1}
11.512(1+2)=32\frac{1}{2}(1 + 2) = \frac{3}{2}8.6×10−28.6 \times 10^{-2}0.50
21.416666…12(32+43)=1712\frac{1}{2}(\frac{3}{2} + \frac{4}{3}) = \frac{17}{12}2.5×10−32.5 \times 10^{-3}0.33
31.4142157577408\frac{577}{408}2.1×10−62.1 \times 10^{-6}0.353
41.41421356237469665857470832\frac{665857}{470832}1.6×10−121.6 \times 10^{-12}0.354
51.4142135623730951below one float spacing

The correct digits go 0, 1, 2, 5, 11, 16. Update 5 moves by 1.6×10−121.6 \times 10^{-12}, more than tol⋅∣x∣=1.4×10−12\mathrm{tol} \cdot |x| = 1.4 \times 10^{-12}, so the method takes update 6, which moves by one unit in the last place, and stops: newton returns (x6,6)(x_6, 6). test_hand_example checks the fractions, the count, and the five points where ff was evaluated.

1/4=0.51/\sqrt{4} = 0.5 from y0=0.4y_0 = 0.4 (δ0=−0.2\delta_0 = -0.2):

y1=0.4(1.5−0.5⋅4⋅0.16)=0.4⋅1.18=0.472,δ1=−0.056,y_1 = 0.4 \left(1.5 - 0.5 \cdot 4 \cdot 0.16\right) = 0.4 \cdot 1.18 = 0.472, \qquad \delta_1 = -0.056 ,

y2=0.472(1.5−2⋅0.222784)=0.472⋅1.054432=0.497691904,δ2=−0.0046.y_2 = 0.472 \left(1.5 - 2 \cdot 0.222784\right) = 0.472 \cdot 1.054432 = 0.497691904, \qquad \delta_2 = -0.0046 .

Check with the error formula: −32(0.2)2−12(−0.2)3=−0.06+0.004=−0.056-\frac{3}{2}(0.2)^2 - \frac{1}{2}(-0.2)^3 = -0.06 + 0.004 = -0.056. These are test_rsqrt_hand_example.

def newton(f, df, x0: float, tol: float = 1e-12, max_iter: int = 50) -> tuple[float, int]: ...
def rsqrt_newton(x: ArrayLike, y0: ArrayLike, iters: int) -> NDArray: ...

newton returns (root, updates). It raises ValueError for a non-finite x0, tol <= 0, or max_iter < 1, and RuntimeError naming the cause when the derivative is zero or not finite (“derivative”), an iterate is not finite (“diverged”), or max_iter updates do not converge (“no convergence”). rsqrt_newton broadcasts x against y0 and computes in their floating dtype (float64 for integers).

TestKINDChecksWhy it matters downstream
test_hand_exampleunit, smokesection 3: the fractions, 6 updates, 5 evaluationsyou and the tests agree on the iteration and the count
test_digits_double_per_steppropertyen+1/en2→0.354e_{n+1}/e_n^2 \to 0.354quadratic convergence, the reason Newton is used at all
test_returns_on_exact_rootboundaryf(x0)=0f(x_0) = 0 returns (x0,0)(x_0, 0)no needless step or miscount
test_relative_tolerance_for_large_rootsboundary1.6×108\sqrt{1.6 \times 10^8} and 3×10203\sqrt[3]{3 \times 10^{20}} convergeroots of every scale
test_converges_on_many_rootsgolden30 seeded square and cube roots to full precisionthe closed forms are the oracle
test_numeric_derivative_when_df_is_noneunitex=3e^x = 3 with df=None gives ln⁡3\ln 3your M01.1 slope inside Newton
test_zero_derivative_raisesboundarya flat or nan tangent raises RuntimeError (“derivative”)failures are reported, not divided by
test_cycle_raises_after_max_iterboundarythe 0, 1, 0, 1 cycle raises RuntimeErrorlocal convergence has a budget
test_max_iter_counts_updatesboundary6 updates fit max_iter = 6, not 5the budget is exact
test_divergence_raisesboundaryx3\sqrt[3]{x} doubles away until inf; RuntimeError (“diverged”)runaway iterates are caught
test_rejects_bad_argumentsboundarynan or inf start, tol <= 0, max_iter = 0caller bugs surface at once
test_rsqrt_hand_exampleunit, smoke0.472, then 0.4977the update formula
test_rsqrt_error_squares_each_stepgoldenfrom the bit-trick guess, float32 errors 3.5e-2, 1.8e-3, 5e-6M09.5’s accuracy target
test_rsqrt_float64_reaches_full_precisiongoldenfour float64 steps reach rounding levelthe error keeps squaring
test_rsqrt_keeps_dtype_and_broadcastsunitfloat32 step for step, bit for bit; ints in float64; iters = 0; iters < 0 raisesthe C port is compared bit for bit
PitfallSymptomCaught by
1. not checking for a zero or non-finite derivativeZeroDivisionError, or a nan that later looks like “no convergence”test_zero_derivative_raises (mutant s07)
2. returning the last iterate when the method did not convergethe cycle’s last point (0.0) reported as a root of x3−2x+2x^3 - 2x + 2test_cycle_raises_after_max_iter (mutant s08)
2b. letting an inf iterate continuethe failure is blamed on a flat tangent one step latertest_divergence_raises (mutant s09)
3. an absolute stopping testno convergence at the right answer near 1.3×1041.3 \times 10^4test_relative_tolerance_for_large_roots (mutant s05)
4. computing the derivative once at x0x_0 (chord method)linear convergence, or divergence when f′f' changes a lottest_digits_double_per_step (mutant s03), test_numeric_derivative_when_df_is_none (mutant s06)
5. dropping a factor of yy in y(32−12xy2)y(\frac{3}{2} - \frac{1}{2} x y^2)converges to the wrong value, 1/x1/x-liketest_rsqrt_hand_example (mutant s10)
6. computing float32 inputs in float64results that differ from the C port in the last bits, and a float64 array where float32 was expectedtest_rsqrt_keeps_dtype_and_broadcasts (mutant s12)
stepping uphill, x+f/f′x + f/f', or multiplying by f′f'the hand example’s iterates go wrong at step 1test_hand_example (mutants s01, s02)
always stepping once, even at an exact root(3.0,1)(3.0, 1) instead of (3.0,0)(3.0, 0)test_returns_on_exact_root (mutant s04)
one Newton step fewer than askedthe error after two steps is 1.8×10−31.8 \times 10^{-3}, not 5×10−65 \times 10^{-6}test_rsqrt_error_squares_each_step (mutant s11)
DirectionModuleHow it uses this
BackM01.1newton(f, None, x0) calls central_diff at every iterate
ForwardM09.5the same y(32−12xy2)y(\frac{3}{2} - \frac{1}{2}xy^2) in C, float32, from the bit-trick guess; its differential test compares against your rsqrt_newton (Pass 6)
ForwardL9.6tl_rmsnorm_f32 multiplies by that reciprocal square root for every token
ForwardM10.6optional: Muon orthogonalizes momentum with Newton-Schulz, Newton’s method for a matrix function
ForwardM09.3convergence order and error propagation, generalized
ForwardM07.7fit_temperature calls newton(g, dg, 0.0) on the derivative of the NLL in 1/T1/T, with the variance of the logits as dg

M07.7 is the registered call site; M09.5 is an optional C side quest and M10.6 is not authored yet.

Your pieceProduction equivalentWhat it addsWhere to look
rsqrt_newton with the bit trickx86 rsqrtps / ARM frsqrte plus one Newton step; CUDA rsqrtfa hardware table lookup for the first 12 bits, then the same refinementIntel intrinsics guide _mm_rsqrt_ps; ggml’s ggml_vec_* RMSNorm paths
newtonSciPy scipy.optimize.newton, brentqsecant and Halley variants; Brent’s method brackets the root so it can never divergescipy/optimize/_zeros_py.py
a Newton step that can failsafeguarded Newton (rtsafe)falls back to bisection whenever the Newton step leaves the bracketPress et al., Numerical Recipes, section 9.4
scalar NewtonNewton-Schulz in MuonNewton’s method for the matrix sign function, X←32X−12XX⊤XX \leftarrow \frac{3}{2}X - \frac{1}{2}XX^\top X, the same polynomial as section 2.4Jordan et al., Muon (2024); M10.6