# Precision and numerics

`defrun :precision` is `:f32` or `:bf16`. Under `:bf16` the compiler
rewrites the graphs: activations and matmuls run in `bf16`, trained
parameters stay `f32` master weights, and nodes annotated as numerically
sensitive run in `f32`. The file states what is sensitive; the compiler does the casting.

## with-numerics

```lisp
(defun rms-norm (x &key (eps 1e-6))
  (* x
     (with-numerics (:sensitivity :high)
       (rsqrt (+ (mean (* x x) :axes [-1] :keepdims true) eps)))))
```

`(with-numerics (:sensitivity :high) body...)` annotates every node created
while the body is evaluated. Under `:bf16` such a node computes in `f32` and
its output stays `f32`, an `f32` island until the value meets a `bf16`
tensor again. Above, the mean of squares and its reciprocal square root are
`f32` while the normalized activation comes back as `bf16`, as PyTorch's
`F.rms_norm` returns its input's dtype. Under `:f32` the annotation changes
nothing.

## matmul into an island

A `bf16` `matmul` rounds its product to `bf16` on store, and a sensitive
node that reads it reads that rounded value, as PyTorch's autocast does.
When the island needs the sum before rounding, ask for it:

```lisp
(let ((v (matmul x w :out :f32)))               ; bf16 operands, f32 result
  (with-numerics (:sensitivity :high) (exp-map v)))
```

`:out :f32` keeps the `bf16` operands and the `bf16` GEMM's speed and
stores its `f32` accumulator, which is an island like `(cast x :f32)`. A
matmul inside the island instead casts its operands up and runs an `f32`
GEMM, which is several times slower. The gradient reaches the operands as
`bf16`, so the backward products stay `bf16` GEMMs. On `f32` operands
`:out :f32` changes no value. The interpreter and Metal hold `bf16` values in
`f32` and round only at a `cast`, so there it is the product they already
compute; CUDA runs cuBLAS's `bf16` GEMM with an `f32` result.

## bf16 regions

```lisp
(defun project (x w)
  (with-numerics (:precision :bf16)
    (matmul x w)))
```

`(with-numerics (:precision :bf16) body...)` is the reverse: every node
created in the body computes in `bf16` under either policy. Its float
inputs are cast down, a parameter's once, as under `:bf16`; its products
and stores round to `bf16`; and under `:f32` its output is cast back up
where a node outside the region reads it. A run can stay `:f32` and put its
products on tensor cores one region at a time.

The innermost form wins. A `(:sensitivity :high)` inside a region is an
`f32` island in it, which is how the library's marks keep `softmax` in
`f32` when a region calls it; a region inside a sensitive form clears the
sensitivity for its own body. A region computes in `bf16` even where every
value it reads is `f32`, so a region around a whole model is not quite the
`:bf16` policy: that keeps a value `f32` until it meets a parameter.

A region adds no choice that can vary between runs. Its dtypes are fixed
when the file compiles and its rounding is the one cast, nearest-even on
every device, so a run with regions keeps the determinism level it has
without them (one device, same bits; determinism).

## gradient

`(gradient y x)` ([tensors](https://sx.041.io/docs/tensors.md#gradient)) appends its backward
nodes while the graph is traced, before the precision policy runs, so
under `:bf16` they are lowered like the forward: an `f32` sweep whose ops
that compute a value each round on store with a `cast` after them. The
training backward is the other way round, differentiated from the lowered
graph, so the two need not agree bit for bit. Each backward node carries
the marks of the forward node it differentiates, so the backward of a
sensitive node is an `f32` island too, and a `with-numerics` around the
call stamps its marks on every node the call appends, over those:

```lisp
(with-numerics (:sensitivity :high) (gradient energy x))
```

computes the whole sweep in `f32` under `:bf16`, though `energy` was not.

## Parameters

`(defparam w init :numerics :high)` keeps a parameter's uses in `f32`: the
master weight is read directly instead of through a cast down. See
[models](https://sx.041.io/docs/models.md#defparam).

A frozen parameter (`:trainable false`) has no master weight to keep. It is
stored in the dtype its readers read: when every use of an `f32` one is
cast down to `bf16`, it is stored as `bf16`, rounded once when it is made
or loaded, which gives those uses the same bits from half the memory. Under
`:bf16` that is any frozen parameter no sensitive node and no `:numerics
:high` reads; under `:f32`, one read only inside `(with-numerics
(:precision :bf16))` regions, such as a teacher whose forward pass is one
region. A frozen parameter read in `f32` anywhere stays `f32`.
[`explain`](https://sx.041.io/docs/explain.md) counts each parameter at the size it is stored, and
so does [memory admission](https://sx.041.io/docs/devices.md#memory).

## What the standard library marks

`logsumexp`, `softmax`, `cross-entropy`, `softcap`, the reciprocal square
root of `rms-norm` and the statistics of `batch-norm` in [sexpgpu/nn](https://sx.041.io/docs/nn.md) carry
`:sensitivity :high`, so a `:bf16` run computes them in `f32` without the
file saying anything.

## Seeing it

[`explain`](https://sx.041.io/docs/explain.md) reports where the precision policy keeps `f32`,
and under `:f32` the lines of its `bf16` regions.
`matrix-scan`, and `linear-scan` through it, also compute in `f32` under
`:bf16`. Use `cast` for explicit conversions; a cast to the same dtype is
the identity. `newton-schulz` in
[sexpgpu/optim](https://sx.041.io/docs/optim.md) runs its iteration in `bf16` under either
precision and returns `f32`.

What else decides a run's bits, and when two runs agree bit for bit, is
[determinism](https://sx.041.io/docs/determinism.md).

Related: [defrun](https://sx.041.io/docs/defrun.md), [tensor operations](https://sx.041.io/docs/tensors.md).

---

SexpGPU documentation. Every page: https://sx.041.io/llms.txt
