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; (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)—batched; leading dims broadcast, inner dims must agree
(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
(linear-scan a b :axis 0 :reverse false):axis :reverseh[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 and b are floats of one dtype and broadcast to one shape; computed in f32 under :bf16. cumsum in sexpgpu/nn is built on it
(cast x :bf16)—same shape; a cast to the same dtype is the identity
(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:std 1.0 :mean 0.0 :f32
(uniform [..]):low :high :dtype:low 0.0 :high 1.0 :f32
(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.

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.

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