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