S-exp GPU
All pages
Docs · Models and optimizersMarkdown

sexpgpu/nn

The standard neural network module, compiled into the binary: (require "sexpgpu/nn" name ...). Every definition is written in SexpGPU itself, so your own models and functions use the same forms; see modules to write your own library.

Functions

definitionsignature and meaning
mean(mean x &key (axes nil) (keepdims false)); every axis by default
logsumexp(logsumexp x &key (axis -1)), max-shifted
softmax(softmax x &key (axis -1))
rms-norm(rms-norm x &key (eps 1e-6)), no gain
rms(rms x), the root mean square over every element, a scalar
relu(relu x)
relu-squared(relu-squared x)
gelu(gelu x), the tanh approximation
softcap(softcap x cap), cap * x * rsqrt(x^2 + cap^2): a smooth clamp to ±cap
cross-entropy(cross-entropy logits targets &key (reduction :mean)); logits [... classes], integer targets [...]; :sum or :mean over every leading position
causal-mask(causal-mask n), [n n] boolean, true where key position <= query position
causal-attention(causal-attention q k v &key (scale nil) (layer nil)); all [batch heads seq head-dim]; scale defaults to 1/sqrt(head-dim)
rotary-frequencies(rotary-frequencies head-dim): (1/1024)^linspace(0,1,head-dim/4) then as many zeros, so half the head dimension does not rotate
rotary(rotary x), x is [batch seq heads head-dim], positions along axis 1
pad(pad x axis before after), zeros in front of and behind x along one axis
conv2d(conv2d x w &key (stride 1) (padding 0)); x is [batch height width in], w is [kh kw in out]; the correlation F.conv2d computes, as one product over the strided windows of the padded input (im2col)
max-pool2d(max-pool2d x &key (size 2)), the largest of every size by size block of [batch height width channels]
batch-norm(batch-norm x &key (eps 1e-5)), every channel (the last axis) centred and scaled by its mean and biased variance over every other axis: batch statistics, no running estimates
argmax(argmax x &key (axis -1)), the :i32 position of the largest entry along axis, the first of equal ones
one-hot(one-hot ids n), integer ids [...] as [... n] with a one at each id
sigmoid(sigmoid x), computed from exp(-|x|) so it never overflows
silu(silu x), x * sigmoid(x)
softplus(softplus x), log(1 + exp(x)) without overflow
causal-conv1d(causal-conv1d x w); x is [batch seq channels], w is [k channels]: each channel convolved with its own k taps over its current and k - 1 previous positions, as F.conv1d(groups=channels, padding=k-1) cut to seq; tap k - 1 weighs the current position
cumsum(cumsum x &key (axis 0) (reverse false)), (linear-scan (ones-like x) x ...)

logsumexp, softmax, cross-entropy, softcap, the reciprocal square root of rms-norm and the statistics of batch-norm compute in f32 under :bf16; see numerics.

Models

constructorsignatureparameters
linear(linear in out &key (bias true) (init-std nil))weight [in out], normal with std init-std or 1/sqrt(in); b [out], zeros
rmsnorm(rmsnorm dim)gain [dim], ones
embedding(embedding vocab dim)table [vocab dim], standard normal
attention(attention dim heads &key (layer nil))q k v proj, each (linear dim dim), proj zero-initialized
mlp(mlp dim &key (hidden nil))fc to hidden (default 4 * dim), relu-squared, proj back, zero-initialized
block(block dim heads &key (layer nil))pre-norm residual: norm1, attn, norm2, ffn
gpt(gpt &key vocab layers dim heads)embed, norm-in, blocks.<i>, norm-out, head (no bias, zero-initialized)

gpt maps integer token ids [batch seq] to logits [batch seq vocab]. Its paths are model.embed.table, model.blocks.3.attn.q.weight, model.head.weight and so on; select them for optimizer groups.

Built-in diagnostics

Off until a run selects them; see diagnostics.

namereducerfrom
attention/entropymeancausal-attention with :layer, so every gpt layer
attention/max-probabilitymaxthe same
residual/rmsmeanblock with :layer, after each block
logits/max-absmaxgpt
logits/rmsmeangpt
logits/grad-rmsmeangpt, training graph only, via tap-gradient

Example: block

block, the unit gpt stacks:

(defmodel block (dim heads &key (layer nil))
  (let ((norm1 (rmsnorm dim))
        (attn (attention dim heads :layer layer))
        (norm2 (rmsnorm dim))
        (ffn (mlp dim)))
    (lambda (x)
      (let* ((x (+ x (attn (norm1 x)))) (out (+ x (ffn (norm2 x)))))
        (unless (null layer)
          (diagnostic "residual/rms" (rms out) :reduce :mean :layer layer))
        out))))

Related: models, sexpgpu/optim.