# 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](https://sx.041.io/docs/how-it-works.md#where-graphs-come-from).

## 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`](https://sx.041.io/docs/explain.md#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

| builtin | keywords | result |
|---|---|---|
| `(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` `:keepdims` | every 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` `:keepdims` | the 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)` | `:out` | batched; 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](https://sx.041.io/docs/numerics.md#matmul-into-an-island) |
| `(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)` | `:axis` | take along one axis; equal ranks, the result has `idx`'s shape |
| `(scatter-add base idx values :axis -1)` | `:axis` | `base` 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` `:projective` | `h[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](#matrix-scan) |
| `(top-k x k :axis -1)` | `:axis` | two values: the `k` largest of the float `x` along `:axis`, largest first, and their positions as `:i32`; see [top-k](#top-k) |
| `(table-scan table symbols :axis -1 :start 0)` | `:axis` `:start` | the 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](#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](#gradient) |
| `(iota [n n] :axis 1 :dtype :i32)` | `:axis` `:dtype` | indices along an axis; `:axis 0`, `:i32` by default; integer dtypes only |
| `(zeros [..])` `(ones [..])` | `:dtype` | `:f32` by default |
| `(full [..] value)` | `:dtype` | one 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](https://sx.041.io/docs/curriculum.md#ctx):

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

Each salt is its own stream; see [determinism](https://sx.041.io/docs/determinism.md#random-draws).
`: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](https://sx.041.io/docs/metrics.md) and [writing diagnostics](https://sx.041.io/docs/writing-diagnostics.md).
Higher-level functions (`mean`, `softmax`, `rms-norm`, `cross-entropy`,
attention) are in [sexpgpu/nn](https://sx.041.io/docs/nn.md).

## top-k

```lisp
(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

```lisp
(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

```lisp
(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](https://sx.041.io/docs/nn.md) 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

```lisp
(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](https://sx.041.io/docs/numerics.md#gradient).

## Errors

| code | when |
|---|---|
| `E-DIM-001` | shapes 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-002` | dtypes differ, or a keyword is not a dtype |
| `E-DIM-003` | a tensor was expected |
| `E-FLOW-001` | `if` on a tensor; use `where`. See [staging](https://sx.041.io/docs/control.md#truth-and-staging) |
| `E-FLOW-002` | a tensor from another graph; recompute it here |
| `E-FLOW-003` | a tensor operation with no graph being traced |
| `E-GRAD-001` | `gradient` of a `y` that is not one float value |
| `E-GRAD-002` | `gradient` with respect to an integer or boolean `x` |
| `E-GRAD-003` | `gradient` of a `y` not computed from `x` |
| `E-GRAD-004` | `gradient` of a `y` computed from a `tap-gradient` lambda's gradient |

```console
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](https://sx.041.io/docs/numerics.md), [models](https://sx.041.io/docs/models.md).

---

SexpGPU documentation. Every page: https://sx.041.io/llms.txt
