S-exp GPU
All pages
Docs · The languageMarkdown

Tensor operations

A tensor is a symbolic value: an operation on tensors records a node in the graph being traced and returns a new tensor with a known shape and dtype. Tensors exist only inside a traced function; see how it works.

Shapes, dtypes, broadcasting

  • Shapes are vectors of non-negative integers, fixed at compile time.
  • Dtypes are :f32 :bf16 :i32 :i64 :bool; anything else is E-DIM-002.
  • An axis may be negative and counts from the end.
  • Broadcasting follows NumPy: right aligned, each dimension equal or 1. Every broadcast appears as an explicit broadcast node in sexpgpu ir.
  • + - * / pow sqrt exp log max and the six comparisons are overloaded: the tensor reading wins when any argument is a tensor, and a number mixed with a tensor becomes a broadcast constant. Plain numbers at the front of an n-ary * fold first, so (* lr scale tensor) is one constant times the tensor.

Operations

builtinkeywordsresult
(neg x) (rsqrt x) (sin x) (cos x) (tanh x) (abs x)—tensor-only unaries; neg and abs also take integers
(acos x) (asin x) (atan2 y x)—angles in radians; acos and asin are NaN outside [-1, 1]; atan2 broadcasts and returns the angle of (x, y) in [-pi, pi]
(maximum a b) (minimum a b)—elementwise pair, broadcast
(where mask a b)—mask is :bool; all three broadcast to one shape
(sum x :axes [-1] :keepdims false):axes :keepdimsevery axis by default; an axis listed twice, also as 1 and -1, is E-DIM-001; (sum xs) over a list of numbers adds them
(max x :axes [..] :keepdims ..):axes :keepdimsthe same, maximum
(reshape x [..])—same element count and dtype
(broadcast-to x [..])—explicit broadcast
(transpose x [0 2 1 3])—a full permutation
(slice x axis start stop)—non-negative integer bounds, start < stop <= dim
(concat axis [a b ...])—equal ranks and equal dims off axis
(matmul a b :out :f32):outbatched; leading dims broadcast, inner dims must agree. :out :f32 on bf16 operands stores the f32 accumulator of the bf16 product instead of rounding it, at bf16 GEMM speed; f32 is the only :out. See numerics
(index table ids)—rows of [n, rest...] by an integer tensor, giving [ids..., rest...]
(index-add base ids values)—base with each row of values added at the row its id names, repeated ids accumulating: the adjoint of index, so values is [ids..., rest...]
(gather x idx :axis -1):axistake along one axis; equal ranks, the result has idx's shape
(scatter-add base idx values :axis -1):axisbase with values added along one axis at idx, repeated positions accumulating: the adjoint of gather; idx and values have one shape
(matrix-scan a b :axis 0 :reverse false :projective false):axis :reverse :projectiveh[t] = a[t] @ h[t-1] + b[t] along :axis from h = 0, so h[0] = b[0]; :reverse true runs from the end. a is [.., n, n] and b [.., n, m], floats of one dtype, n at most 16; their leading dims broadcast to one shape and :axis counts among them, so -1 is the one before the matrices. The result has b's shape; computed in f32 under :bf16. See matrix-scan
(top-k x k :axis -1):axistwo values: the k largest of the float x along :axis, largest first, and their positions as :i32; see top-k
(table-scan table symbols :axis -1 :start 0):axis :startthe states of the automaton table ([states symbols] of :i32) run over the :i32 symbols along :axis from state :start: h[t] = table[h[t-1], symbols[t]], in symbols' shape; no gradient. See table-scan
(cast x :bf16)—same shape; a cast to the same dtype is the identity
(gradient y x)—the gradient of the float scalar y with respect to the float tensor x it is computed from, both traced in the current graph, appended as ordinary nodes; in any graph, so a generator or an evaluation pass can ascend on its own inputs. See gradient
(iota [n n] :axis 1 :dtype :i32):axis :dtypeindices along an axis; :axis 0, :i32 by default; integer dtypes only
(zeros [..]) (ones [..]):dtype:f32 by default
(full [..] value):dtypeone repeated value
(zeros-like x) (ones-like x)—x's shape and dtype
(normal [..]):std :mean :dtype :salt:std 1.0 :mean 0.0 :f32, unsalted
(uniform [..]):low :high :dtype :salt:low 0.0 :high 1.0 :f32, unsalted
(shape x) (rank x) (dim x axis) (numel x) (dtype x)—ordinary values, for compile-time decisions

A degenerate spread (:std 0.0, or :low equal to :high) emits a constant instead of a random node, so a zero-initialized projection costs no random stream. Random streams fold in defrun :seed.

A draw is the same in every call of its graph: every step, every microbatch. :salt makes it new wherever the salt moves. The salt is one integer, a number or an integer scalar tensor that the draw reads when the graph runs, such as (+ (* (getf ctx :step) 65536) (getf ctx :microbatch)) in a graph that is given ctx:

(uniform [4] :low -1.0 :high 1.0 :salt (getf ctx :step))

Each salt is its own stream; see determinism. :salt 0 is a different stream from no salt. A float salt is E-DIM-002, as is a literal salt beyond 2^53 in magnitude, which a constant cannot hold exactly; a salt with more than one element is E-DIM-001.

Reporting operations (metric, diagnostic, tap-gradient, counter) are in metrics and writing diagnostics. Higher-level functions (mean, softmax, rms-norm, cross-entropy, attention) are in sexpgpu/nn.

top-k

(defun keep-best (pool score)                 ; pool [n d], score [n]
  (multiple-value-bind (best at) (top-k score 32 :axis 0)
    (values (index pool at) best)))            ; [32 d] rows, best first

(top-k x k :axis a) returns two values, both x's shape with axis a cut to k: the k largest values along it and their positions as :i32. k is an integer from 1 to the axis's length, else E-DIM-001, as is a rank-0 x, which has no axis to select along. The order is the same on every device: larger first, every NaN above infinity, -0.0 equal to 0.0, and equal values by the lower position, so the positions agree bit for bit with the interpreter. k equal to the axis sorts it.

The values are (gather x at :axis a), so their gradient reaches only the chosen elements, and the positions have none. Taking rows by them, with index for the first axis or gather along any other, keeps the rows a model then reads: its backward runs over those k rows, where a mask over the pool would run it over every row.

table-scan

(defun parity (bits)                          ; bits [b t] :i32 of 0 and 1
  (let* ((s (+ (iota [2 2] :axis 0) (iota [2 2] :axis 1)))
         (flip (where (>= s 2) (- s 2) s)))     ; [[0 1] [1 0]]
    (table-scan flip bits)))                    ; [b t]: parity so far

(table-scan table symbols :axis a :start s) runs a deterministic automaton along axis a of symbols, the last by default: the state after each symbol is table[previous, symbol], starting from state s, 0 by default, and the result holds every state, symbols' shape. table is [states symbols], both tensors :i32, else E-DIM-002; a table of another rank or with no symbols, an axis out of range, or a :start that is not a state is E-DIM-001. The states are integers and have no gradient; a model reads them through index or gather, whose values carry it. A symbol outside [0, symbols), or a table entry the run reaches outside [0, states), is refused by the CPU and Metal, as gather refuses an index, and clamped into range on CUDA, as gather is there.

Each step reads the state before it, so a run of t symbols is t dependent lookups however it is written; table-scan makes it one node and, on CUDA, one kernel with a thread per sequence and the table in shared memory when it fits.

matrix-scan

(defun delta-rule (k v beta)        ; k [b t n], v [b t m], beta [b t 1 1]
  (let* ((column (reshape k (append (shape k) [1])))           ; [b t n 1]
         (row (reshape k [(dim k 0) (dim k 1) 1 (dim k 2)]))   ; [b t 1 n]
         (eye (cast (= (iota [(dim k 2) (dim k 2)] :axis 0)
                       (iota [(dim k 2) (dim k 2)] :axis 1))
                    :f32))
         (write (matmul column (reshape v [(dim v 0) (dim v 1) 1 (dim v 2)]))))
    (matrix-scan (- eye (* beta (matmul column row))) (* beta write) :axis 1)))

(matrix-scan a b :axis t) carries an n by m state through the first-order recurrence h[t] = a[t] @ h[t-1] + b[t], one matrix product a step; above, the delta rule S[t] = (I - beta k k^T) S[t-1] + beta k v^T over [b t n m]. Every column of h is its own recurrence on the shared a. Its gradient is the same scan run backwards on a's transposes. The prelude's linear-scan is this over one-by-one matrices, which is how (linear-scan a b :axis 0 :reverse false) takes a and b, numbers included, broadcast to one shape, and cumsum in sexpgpu/nn is built on that.

:projective true is for a product of projective maps, such as a carried Möbius transformation as a real matrix, whose entries would overflow f32: the running product of a, from the identity, is divided by its largest magnitude after every step, and h by the same number. The result is h up to a positive factor per step when b is zero after its first step, and has no gradient: differentiating through it is an error.

gradient

(defun prepare (batch)
  (let* ((x (field batch :x))
         (energy (sum (* x x x))))
    (list :inputs (+ x (* 0.1 (gradient energy x))) :targets x)))

(gradient y x) runs the backward sweep from y to x while the graph is traced and appends its nodes after what the graph holds so far, so the result is an ordinary tensor of x's shape and dtype that later code reads like any other. y must be one float value and x a float tensor y is computed from; anything else is refused rather than answered with zeros (E-GRAD-001 to E-GRAD-003). A y computed from the gradient a tap-gradient lambda is given is E-GRAD-004: that gradient is filled in after the training backward, so nothing differentiates through it.

Training differentiates through every gradient in the training graph, as through any other node. A prepare or model that moves its inputs uphill on the student's own score therefore adds second-order terms to the parameter gradient, and the step pays for a second backward through the first. There is no stop-gradient form yet that would make the ascent a constant to training; that is a known gap.

Each backward node carries the numerics marks of the forward node it differentiates, and a with-numerics around the call stamps its own over them on every node the call appends. Under :bf16 the sweep is built on the graph as traced and precision lowering then rewrites its nodes with the rest, so each op that computes a value rounds on store; the training backward is differentiated after lowering instead. See numerics.

Errors

codewhen
E-DIM-001shapes do not fit; the notes name both operands as you wrote them, with shape and dtype, and any broadcast that would have applied
E-DIM-002dtypes differ, or a keyword is not a dtype
E-DIM-003a tensor was expected
E-FLOW-001if on a tensor; use where. See staging
E-FLOW-002a tensor from another graph; recompute it here
E-FLOW-003a tensor operation with no graph being traced
E-GRAD-001gradient of a y that is not one float value
E-GRAD-002gradient with respect to an integer or boolean x
E-GRAD-003gradient of a y not computed from x
E-GRAD-004gradient of a y computed from a tap-gradient lambda's gradient
error[E-DIM-001]: matmul: inner dims differ: [2, 16] x [32, 8]
  --> bad-shape.sx:7:15
   |
 7 |   (lambda (x) (matmul x weight)))
   |               ^^^^^^^^^^^^^^^^^
note: operands: `x` [2, 16] f32, `weight` [32, 8] f32

Related: precision and numerics, models.