S-exp GPU
All pages
Docs · Models and optimizersMarkdown

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

(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:

(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

(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) 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:

(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.

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 counts each parameter at the size it is stored, and so does memory admission.

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 carry :sensitivity :high, so a :bf16 run computes them in f32 without the file saying anything.

Seeing it

explain 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 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.

Related: defrun, tensor operations.