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

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.

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. linear-scan also computes 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.

Related: defrun, tensor operations.