All pages
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.