# tinygrad/tensor.bend -- the lazy tensor core of tinygrad-bend. # # Port of upstream tinygrad's semantic core: # - Tensor is LAZY: ops only extend a pending graph (tensor.py `_apply_uop`, # engine/lazy.py); nothing computes until `realize` (here: `eval`). # - elementwise ops broadcast like `_broadcasted` (align right, 1s stretch) # - movement ops are pure views (mixin/movement.py) # - reductions drop the reduced axis (mixin/reduce.py `_reduce`) # - matmul is batched dot ((...,m,k) @ (...,k,n) -> (...,m,n)), mixin/op.py # # Kernel idiom: Bend outlaws mutual recursion and matching computed values, # so every eval kernel is ONE self-recursive def whose pending Array reads # travel as ordinary arguments, destructured at the top of each step. # # Affine-world deviations from upstream (see README): # - F32 only (Bend has no f64; Metal has none either) # - dimensions are >= 1 (no size-0 tensors) # - a Param's data lives in a caller-owned Params table: sharing a # parameter across a graph needs no Array.clone -- the table is the # buffer, mirroring tinygrad's realized-buffer identity. # - shapes ride along in every node (affinity forbids re-reading a # consumed Tensor to ask for its shape) import Base import ./shape.bend as Sh # ---------- types ---------- # elementwise binary ops. Comparisons answer 1.0 / 0.0 tensors. type BinOp is Data: OpAdd{} OpSub{} OpMul{} OpDiv{} OpPow{} OpMax2{} OpEq{} OpNe{} OpLt{} OpLe{} OpGt{} OpGe{} # elementwise unary ops type UnOp is Data: OpNeg{} OpExp{} OpLog{} OpSqrt{} OpAbs{} OpTanh{} OpSigmoid{} OpRelu{} # movement ops (pure views) type MovOp is Data: OpReshape{} OpPermute{} OpFlip{} OpExpand{} OpPad{} OpShrink{} # reduction ops type RedOp is Data: OpSum{} OpMax{} OpMean{} # fused backward kernel ops: unary (a, g) and binary (a, b, g) type Bwd1Op is Data: BSigmoid{} BTanh{} BSqrt{} BExp{} BNeg{} BLog{} BAbs{} BRelu{} BId1{} type Bwd2Op is Data: BDivL{} BDivR{} BPowL{} BPowR{} BMax2L{} BMax2R{} BId{} B2Neg{} BMulL{} BMulR{} BZero{} BTriL{} BTriR{} # One pending node of the computation graph. Affine (is Type): every op # consumes its inputs, exactly like a UOp in upstream tinygrad. type Tensor is Type: # a realized constant Leaf{buf: Array, shape: List<&2, Nat>} # a reference into the caller-owned parameter table (grad tracked) Param{pid: Nat, off: Nat, shape: List<&2, Nat>} # broadcast elementwise binary: sa/sb are the PRE-broadcast shapes Bin{op: BinOp, a: Tensor, b: Tensor, sa: List<&2, Nat>, sb: List<&2, Nat>, out: List<&2, Nat>} # elementwise unary Un{op: UnOp, a: Tensor, out: List<&2, Nat>} # ternary where(mask 1.0/0.0, then, else) Tri{c: Tensor, a: Tensor, b: Tensor, sc: List<&2, Nat>, sa: List<&2, Nat>, sb: List<&2, Nat>, out: List<&2, Nat>} # batched matmul over equal (already broadcast) leading dims Mat{a: Tensor, b: Tensor, batch: List<&2, Nat>, m: Nat, k: Nat, n: Nat, out: List<&2, Nat>} # movement view: arg1/arg2 are per-op args; sin is the source shape Mov{op: MovOp, a: Tensor, arg1: List<&2, Nat>, arg2: List<&2, Nat>, sin: List<&2, Nat>, out: List<&2, Nat>} # reduction over ONE axis (drop the axis; keepdim via explicit reshape) Red{op: RedOp, a: Tensor, axis: Nat, sin: List<&2, Nat>, out: List<&2, Nat>} # fused backward kernels (VJP helpers whose formula reuses a value, which # affinity forbids at the Tensor level; F32 locals are freely reusable). # These only ever appear inside backward graphs, whose own gradient stops # here (double backward is unsupported -- see UPSTREAM.md). Bwd1{op: Bwd1Op, a: Tensor, g: Tensor, sin: List<&2, Nat>, out: List<&2, Nat>} Bwd2{op: Bwd2Op, a: Tensor, b: Tensor, g: Tensor, sa: List<&2, Nat>, sb: List<&2, Nat>, out: List<&2, Nat>} # ---------- allocation ---------- # depth_of -- ceil(log2(n)): the Array.new depth for n slots. Structural # fuel-peeling with a doubling capacity (Bend rejects Nat.div shrinking). def df_go(+fuel: Nat, +cap: Nat, +n: Nat, +acc: Nat, z: Bool) -> Nat: match fuel cap n acc z: case 0n, _, _, _, _: acc case 1n+f2, _, _, _, False{}: acc case 1n+f2, _, _, _, True{}: df_go(f2, (Nat.add(cap, cap) : Nat), n, (1n+acc : Nat), Nat.is_lt(Nat.add(cap, cap), n)) def depth_for(+n: Nat) -> Nat: df_go(n, 1n, n, 0n, Nat.is_lt(1n, n)) def alloc(+n: Nat) -> Array: Array.new(F32, depth_for(n), 0.0) # ---------- creation (tinygrad Tensor.full / zeros / ones / uniform) ---------- def fill_leaf(rem: Nat, +idx: Nat, +v: F32, a: Array, +shape: List<&2, Nat>) -> Tensor: match rem: case 0n: Leaf{a, shape} case 1n+p: fill_leaf(p, 1n+idx, v, Array.set(F32, a, U32.from_nat(idx), v), shape) def full_go(+shape: List<&2, Nat>, +v: F32, z: Bool) -> Tensor: match z: case True{}: fill_leaf(1n, 0n, v, alloc(1n), shape) case False{}: fill_leaf(Sh.numel(shape), 0n, v, alloc(Sh.numel(shape)), shape) def Tensor.full(+shape: List<&2, Nat>, +v: F32) -> Tensor: full_go(shape, v, List.is_empty(&2, Nat, shape)) def Tensor.zeros(+shape: List<&2, Nat>) -> Tensor: Tensor.full(shape, 0.0) def Tensor.ones(+shape: List<&2, Nat>) -> Tensor: Tensor.full(shape, 1.0) def Tensor.const(+v: F32) -> Tensor: Tensor.full(1n <> Nil{}, v) # matrix shape helpers def matrix_shape(+r: Nat, +c: Nat) -> List<&2, Nat>: r <> c <> Nil{} def row_width(+rows: List<&2, List<&2, F32>>) -> Nat: match rows: case Nil{}: 0n case Con{h, _}: List.length(&2, F32, h) def concat_rows(+rows: List<&2, List<&2, F32>>) -> List<&2, F32>: match rows: case Nil{}: Nil{} case Con{h, t}: List.append(&2, F32, h, concat_rows(t)) # Tensor.matrix via a flat fill: rows must all have the same width def fill_from_list(rem: Nat, +xs: List<&2, F32>, +idx: Nat, a: Array, +shape: List<&2, Nat>) -> Tensor: match rem xs: case 0n, _: Leaf{a, shape} case 1n+p, Nil{}: Leaf{a, shape} case 1n+p, Con{h, t}: fill_from_list(p, t, 1n+idx, Array.set(F32, a, U32.from_nat(idx), h), shape) def Tensor.matrix(+rows: List<&2, List<&2, F32>>) -> Tensor: flat = concat_rows(rows) +r = List.length(&2, List<&2, F32>, rows) +c = row_width(rows) fill_from_list(r*c, flat, 0n, alloc(r*c), matrix_shape(r, c)) def Tensor.vector(+xs: List<&2, F32>) -> Tensor: +n = List.length(&2, F32, xs) fill_from_list(n, xs, 0n, alloc(n), n <> Nil{}) # seeded uniform rng in [lo, hi), an LCG like tinygrad's Tensor.uniform def lcg_next(+s: U32) -> U32: U32.add(U32.mul(s, 1664525), 1013904223) def lcg_val(+s: U32) -> F32: F32.div(U32.to_f32(U32.shrn(s, 8n)), 16777216.0) def lcg_between(+lo: F32, +hi: F32, +s: U32) -> F32: F32.add(lo, F32.mul(F32.sub(hi, lo), lcg_val(s))) def fill_uniform(rem: Nat, +idx: Nat, +shape: List<&2, Nat>, +lo: F32, +hi: F32, +s: U32, a: Array) -> Tensor: match rem: case 0n: Leaf{a, shape} case 1n+p: fill_uniform(p, 1n+idx, shape, lo, hi, lcg_next(s), Array.set(F32, a, U32.from_nat(idx), lcg_between(lo, hi, s))) def Tensor.uniform(+shape: List<&2, Nat>, +lo: F32, +hi: F32, +seed: U32) -> Tensor: fill_uniform(Sh.numel(shape), 0n, shape, lo, hi, seed, alloc(Sh.numel(shape))) # ---------- parameters ---------- # The parameter table is pure DATA (a flat F32 list), so, unlike an # Array, it can be copied freely to every Param leaf that reads it. # Registry = the shapes of all parameters, ordered by pid. def offset_of(+reg: List<&2, List<&2, Nat>>, +pid: Nat) -> Nat: match reg pid: case _, 0n: 0n case Nil{}, _: 0n case Con{h, t}, 1n+p2: Nat.add(Sh.numel(h), offset_of(t, p2)) def reg_at(+reg: List<&2, List<&2, Nat>>, +pid: Nat) -> List<&2, Nat>: match reg pid: case Nil{}, _: Nil{} case Con{h, _}, 0n: h case Con{_, t}, 1n+p2: reg_at(t, p2) def Tensor.param(+reg: List<&2, List<&2, Nat>>, +pid: Nat) -> Tensor: Param{pid, offset_of(reg, pid), reg_at(reg, pid)} # copy one parameter out of the flat table into a fresh dense array def copy_num(rem: Nat, +xs: List<&2, F32>, +idx: Nat, a: Array) -> Array: match rem xs: case 0n, _: a case 1n+p, Nil{}: a case 1n+p, Con{h, t}: copy_num(p, t, 1n+idx, Array.set(F32, a, U32.from_nat(idx), h)) def params_get(+ps: List<&2, F32>, +off: Nat, +numel: Nat) -> Array: copy_num(numel, List.drop(&2, F32, ps, off), 0n, alloc(numel)) # NOTE on affine use: a Tensor value is used exactly once. Constants that # feed two ops are simply built twice; parameters that feed two ops are two # Param nodes over the same table entry (they share for free, and their # gradients accumulate by pid -- tinygrad's `t.grad + g` semantics). # ---------- graph building: elementwise with broadcast (tensor.py _broadcasted) ---------- def Tensor.bin(+op: BinOp, a: Tensor, b: Tensor, +sa: List<&2, Nat>, +sb: List<&2, Nat>) -> Tensor: Bin{op, a, b, sa, sb, Sh.bcast(sa, sb)} def Tensor.add(a: Tensor, b: Tensor, +sa: List<&2, Nat>, +sb: List<&2, Nat>) -> Tensor: Tensor.bin(OpAdd{}, a, b, sa, sb) def Tensor.sub(a: Tensor, b: Tensor, +sa: List<&2, Nat>, +sb: List<&2, Nat>) -> Tensor: Tensor.bin(OpSub{}, a, b, sa, sb) def Tensor.mul(a: Tensor, b: Tensor, +sa: List<&2, Nat>, +sb: List<&2, Nat>) -> Tensor: Tensor.bin(OpMul{}, a, b, sa, sb) def Tensor.div(a: Tensor, b: Tensor, +sa: List<&2, Nat>, +sb: List<&2, Nat>) -> Tensor: Tensor.bin(OpDiv{}, a, b, sa, sb) def Tensor.pow(a: Tensor, b: Tensor, +sa: List<&2, Nat>, +sb: List<&2, Nat>) -> Tensor: Tensor.bin(OpPow{}, a, b, sa, sb) def Tensor.max2(a: Tensor, b: Tensor, +sa: List<&2, Nat>, +sb: List<&2, Nat>) -> Tensor: Tensor.bin(OpMax2{}, a, b, sa, sb) def Tensor.eq(a: Tensor, b: Tensor, +sa: List<&2, Nat>, +sb: List<&2, Nat>) -> Tensor: Tensor.bin(OpEq{}, a, b, sa, sb) def Tensor.lt(a: Tensor, b: Tensor, +sa: List<&2, Nat>, +sb: List<&2, Nat>) -> Tensor: Tensor.bin(OpLt{}, a, b, sa, sb) def Tensor.gt(a: Tensor, b: Tensor, +sa: List<&2, Nat>, +sb: List<&2, Nat>) -> Tensor: Tensor.bin(OpGt{}, a, b, sa, sb) def Tensor.neg(a: Tensor, +sa: List<&2, Nat>) -> Tensor: Un{OpNeg{}, a, sa} def Tensor.exp(a: Tensor, +sa: List<&2, Nat>) -> Tensor: Un{OpExp{}, a, sa} def Tensor.log(a: Tensor, +sa: List<&2, Nat>) -> Tensor: Un{OpLog{}, a, sa} def Tensor.sqrt(a: Tensor, +sa: List<&2, Nat>) -> Tensor: Un{OpSqrt{}, a, sa} def Tensor.abs(a: Tensor, +sa: List<&2, Nat>) -> Tensor: Un{OpAbs{}, a, sa} def Tensor.tanh(a: Tensor, +sa: List<&2, Nat>) -> Tensor: Un{OpTanh{}, a, sa} def Tensor.sigmoid(a: Tensor, +sa: List<&2, Nat>) -> Tensor: Un{OpSigmoid{}, a, sa} def Tensor.relu(a: Tensor, +sa: List<&2, Nat>) -> Tensor: Un{OpRelu{}, a, sa} def Tensor.where(c: Tensor, a: Tensor, b: Tensor, +sc: List<&2, Nat>, +sa: List<&2, Nat>, +sb: List<&2, Nat>) -> Tensor: Tri{c, a, b, sc, sa, sb, Sh.bcast(sc, Sh.bcast(sa, sb))} # ---------- movement ops (mixin/movement.py) ---------- def permute_shape(+s: List<&2, Nat>, +order: List<&2, Nat>) -> List<&2, Nat>: match order: case Nil{}: Nil{} case Con{h, t}: Sh.at(s, h) <> permute_shape(s, t) def Tensor.reshape(a: Tensor, +sin: List<&2, Nat>, +target: List<&2, Nat>) -> Tensor: Mov{OpReshape{}, a, target, Nil{}, sin, target} def Tensor.permute(a: Tensor, +sin: List<&2, Nat>, +order: List<&2, Nat>) -> Tensor: Mov{OpPermute{}, a, order, Nil{}, sin, permute_shape(sin, order)} def Tensor.flip(a: Tensor, +sin: List<&2, Nat>, +mask: List<&2, Nat>) -> Tensor: Mov{OpFlip{}, a, mask, Nil{}, sin, sin} def Tensor.expand(a: Tensor, +sin: List<&2, Nat>, +target: List<&2, Nat>) -> Tensor: Mov{OpExpand{}, a, target, Nil{}, sin, target} def Tensor.pad(a: Tensor, +sin: List<&2, Nat>, +los: List<&2, Nat>, +his: List<&2, Nat>) -> Tensor: Mov{OpPad{}, a, los, his, sin, Sh.pad_shape(sin, los, his)} def Tensor.shrink(a: Tensor, +sin: List<&2, Nat>, +los: List<&2, Nat>, +his: List<&2, Nat>) -> Tensor: Mov{OpShrink{}, a, los, his, sin, Sh.shrink_shape(sin, los, his)} # ---------- evaluation kernels ---------- # # Kernel idiom: Bend outlaws mutual recursion and matching computed values, # so every kernel is ONE self-recursive def whose pending Array reads travel # as ordinary arguments, destructured at the top of each step. ev dispatches # on the node, recursively evaluates children (self-calls only), then hands # the arrays to a loop; loops never call back into ev. # --- unary: apply then loop --- def apply_un(+op: UnOp, +x: F32) -> F32: match op: case OpNeg{}: F32.neg(x) case OpExp{}: F32.exp(x) case OpLog{}: F32.log(x) case OpSqrt{}: F32.sqrt(x) case OpAbs{}: F32.abs(x) case OpTanh{}: F32.tanh(x) case OpSigmoid{}: F32.div(1.0, F32.add(1.0, F32.exp(F32.neg(x)))) case OpRelu{}: F32.max(x, 0.0) def ew1(+op: UnOp, rem: Nat, +idx: Nat, r: Array & F32, out_a: Array) -> Array: match rem: case 0n: out_a case 1n+p: (s2, x) = r ew1(op, p, 1n+idx, Array.get(F32, s2, U32.from_nat(1n+idx)), Array.set(F32, out_a, U32.from_nat(idx), apply_un(op, x))) # --- binary: apply then loop (broadcast via Sh.bcast_flat) --- def ew2_next(src: Array, +s: List<&2, Nat>, +out: List<&2, Nat>, +idx: Nat) -> Array & F32: Array.get(F32, src, U32.from_nat(Sh.bcast_flat(Sh.unflat(idx, out), s, out))) def apply_bin(+op: BinOp, +x: F32, +y: F32) -> F32: match op: case OpAdd{}: F32.add(x, y) case OpSub{}: F32.sub(x, y) case OpMul{}: F32.mul(x, y) case OpDiv{}: F32.div(x, y) case OpPow{}: F32.pow(x, y) case OpMax2{}: F32.max(x, y) case OpEq{}: F32.from_nat(U32.to_nat(Bool.to_u32(F32.is_eq(x, y)))) case OpNe{}: F32.from_nat(U32.to_nat(Bool.to_u32(F32.is_ne(x, y)))) case OpLt{}: F32.from_nat(U32.to_nat(Bool.to_u32(F32.is_lt(x, y)))) case OpLe{}: F32.from_nat(U32.to_nat(Bool.to_u32(F32.is_le(x, y)))) case OpGt{}: F32.from_nat(U32.to_nat(Bool.to_u32(F32.is_gt(x, y)))) case OpGe{}: F32.from_nat(U32.to_nat(Bool.to_u32(F32.is_ge(x, y)))) def ew2(+op: BinOp, rem: Nat, +idx: Nat, +sa: List<&2, Nat>, +sb: List<&2, Nat>, +out: List<&2, Nat>, ra: Array & F32, rb: Array & F32, out_a: Array) -> Array: match rem: case 0n: out_a case 1n+p: (sa2, x) = ra (sb2, y) = rb ew2(op, p, 1n+idx, sa, sb, out, ew2_next(sa2, sa, out, idx), ew2_next(sb2, sb, out, idx), Array.set(F32, out_a, U32.from_nat(idx), apply_bin(op, x, y))) # --- where: loop over three broadcast sources --- def ew3(rem: Nat, +idx: Nat, +sc: List<&2, Nat>, +sa: List<&2, Nat>, +sb: List<&2, Nat>, +out: List<&2, Nat>, rc: Array & F32, ra: Array & F32, rb: Array & F32, out_a: Array) -> Array: match rem: case 0n: out_a case 1n+p: (sc2, cv) = rc (sa2, av) = ra (sb2, bv) = rb ew3(p, 1n+idx, sc, sa, sb, out, ew2_next(sc2, sc, out, idx), ew2_next(sa2, sa, out, idx), ew2_next(sb2, sb, out, idx), Array.set(F32, out_a, U32.from_nat(idx), Bool.pick(F32, F32.is_gt(cv, 0.5), av, bv))) # --- movement: index mapping then loop --- def permute_mi(+mi: List<&2, Nat>, +order: List<&2, Nat>) -> List<&2, Nat>: match order: case Nil{}: Nil{} case Con{h, t}: Sh.at(mi, h) <> permute_mi(mi, t) def flip_one(+f: Nat, +i: Nat, +d: Nat) -> Nat: match f: case 0n: i case 1n+d2: Nat.sub(Nat.sub(d, 1n), i) def flip_mi(+mi: List<&2, Nat>, +mask: List<&2, Nat>, +sin: List<&2, Nat>) -> List<&2, Nat>: match mi mask sin: case Nil{}, _, _: Nil{} case _, Nil{}, _: mi case _, _, Nil{}: mi case Con{i, it}, Con{f, ft}, Con{d, ds}: flip_one(f, i, d) <> flip_mi(it, ft, ds) def pad_mask_one(+i: Nat, +lo: Nat) -> F32: Bool.pick(F32, Nat.is_lt(i, lo), 0.0, 1.0) def pad_mask(+mi: List<&2, Nat>, +los: List<&2, Nat>) -> F32: match mi los: case Nil{}, _: 1.0 case _, Nil{}: 1.0 case Con{i, it}, Con{lo, lot}: F32.mul(pad_mask(it, lot), pad_mask_one(i, lo)) def pad_src_mi(+mi: List<&2, Nat>, +los: List<&2, Nat>) -> List<&2, Nat>: match mi los: case Nil{}, _: Nil{} case _, Nil{}: mi case Con{i, it}, Con{lo, lot}: Nat.sub(i, lo) <> pad_src_mi(it, lot) def mov_src_idx(+op: MovOp, +idx: Nat, +arg1: List<&2, Nat>, +arg2: List<&2, Nat>, +sin: List<&2, Nat>, +out: List<&2, Nat>) -> Nat: match op: case OpReshape{}: idx case OpPermute{}: Sh.flat(permute_mi(Sh.unflat(idx, out), arg1), sin) case OpFlip{}: Sh.flat(flip_mi(Sh.unflat(idx, out), arg1, sin), sin) case OpExpand{}: Sh.bcast_flat(Sh.unflat(idx, out), sin, out) case OpPad{}: Sh.flat(pad_src_mi(Sh.unflat(idx, out), arg1), sin) case OpShrink{}: Sh.flat(Sh.unflat(idx, out), sin) def mov_mask(+op: MovOp, +idx: Nat, +arg1: List<&2, Nat>, +out: List<&2, Nat>, +x: F32) -> F32: match op: case OpPad{}: F32.mul(x, pad_mask(Sh.unflat(idx, out), arg1)) case _: x def movloop(+op: MovOp, rem: Nat, +idx: Nat, +arg1: List<&2, Nat>, +arg2: List<&2, Nat>, +sin: List<&2, Nat>, +out: List<&2, Nat>, out_a: Array, r: Array & F32) -> Array: match rem: case 0n: out_a case 1n+p: (s2, x) = r movloop(op, p, 1n+idx, arg1, arg2, sin, out, Array.set(F32, out_a, U32.from_nat(idx), mov_mask(op, idx, arg1, out, x)), Array.get(F32, s2, U32.from_nat(mov_src_idx(op, 1n+idx, arg1, arg2, sin, out)))) # --- reduction: stride walk per output slot --- def stride_tail(+s: List<&2, Nat>, +k: Nat) -> List<&2, Nat>: match s k: case Nil{}, _: Nil{} case Con{_, t}, 0n: t case Con{_, t}, 1n+k2: stride_tail(t, k2) def init_acc(+op: RedOp) -> F32: match op: case OpMax{}: F32.neg(3.0e38) case _: 0.0 def insert_at(+mi: List<&2, Nat>, +k: Nat, +v: Nat) -> List<&2, Nat>: match mi k: case _, 0n: v <> mi case Nil{}, _: v <> Nil{} case Con{h, t}, 1n+k2: h <> insert_at(t, k2, v) def red_base(+o: Nat, +out: List<&2, Nat>, +sin: List<&2, Nat>, +axis: Nat) -> Nat: Sh.flat(insert_at(Sh.unflat(o, out), axis, 0n), sin) def red_step(+op: RedOp, +acc: F32, +x: F32) -> F32: match op: case OpSum{}: F32.add(acc, x) case OpMax{}: F32.max(acc, x) case OpMean{}: F32.add(acc, x) def red_done(+op: RedOp, +acc: F32, +dim: Nat) -> F32: match op: case OpSum{}: acc case OpMax{}: acc case OpMean{}: F32.div(acc, F32.from_nat(dim)) def redloop(+op: RedOp, rem_out: Nat, rem_in: Nat, +o: Nat, +j: Nat, +base: Nat, +stride: Nat, +dim: Nat, +acc: F32, +axis: Nat, +sin: List<&2, Nat>, +out: List<&2, Nat>, out_a: Array, r: Array & F32) -> Array: match rem_out rem_in: case 0n, _: out_a case 1n+p, 0n: (s2, _) = r redloop(op, p, dim, 1n+o, 0n, red_base(1n+o, out, sin, axis), stride, dim, init_acc(op), axis, sin, out, Array.set(F32, out_a, U32.from_nat(o), red_done(op, acc, dim)), Array.get(F32, s2, U32.from_nat(red_base(1n+o, out, sin, axis)))) case 1n+p, 1n+q: (s2, x) = r redloop(op, 1n+p, q, o, 1n+j, base, stride, dim, red_step(op, acc, x), axis, sin, out, out_a, Array.get(F32, s2, U32.from_nat(Nat.add(base, Nat.mul(j, stride))))) # --- batched matmul: (batch..., m, k) @ (batch..., k, n) -> (batch..., m, n) --- def mm_aflat(+o: Nat, +m: Nat, +k: Nat, +n: Nat) -> Nat: bi = Nat.div(o, Nat.mul(m, n)) r = Nat.mod(o, Nat.mul(m, n)) i = Nat.div(r, n) Nat.add(Nat.mul(bi, Nat.mul(m, k)), Nat.mul(i, k)) def mm_bflat(+o: Nat, +m: Nat, +k: Nat, +n: Nat) -> Nat: bi = Nat.div(o, Nat.mul(m, n)) r = Nat.mod(o, Nat.mul(m, n)) j = Nat.mod(r, n) Nat.add(Nat.mul(bi, Nat.mul(k, n)), j) def mm_afloat(+o: Nat, +m: Nat, +k: Nat, +n: Nat) -> Nat: r = Nat.mod(o, Nat.mul(m, n)) i = Nat.div(r, n) Nat.mul(i, k) def mm_bfloat(+o: Nat, +m: Nat, +k: Nat, +n: Nat) -> Nat: r = Nat.mod(o, Nat.mul(m, n)) j = Nat.mod(r, n) j def mmloop(rem: Nat, kk: Nat, +o: Nat, +j: Nat, +acc: F32, +m: Nat, +k: Nat, +n: Nat, ra: Array & F32, rb: Array & F32, out_a: Array) -> Array: match rem kk: case 0n, _: out_a case 1n+p, 0n: (sa2, _) = ra (sb2, _) = rb mmloop(p, k, 1n+o, 0n, 0.0, m, k, n, Array.get(F32, sa2, U32.from_nat(Nat.add(mm_aflat(1n+o, m, k, n), mm_afloat(1n+o, m, k, n)))), Array.get(F32, sb2, U32.from_nat(Nat.add(mm_bflat(1n+o, m, k, n), mm_bfloat(1n+o, m, k, n)))), Array.set(F32, out_a, U32.from_nat(o), acc)) case 1n+p, 1n+q: (sa2, av) = ra (sb2, bv) = rb mmloop(1n+p, q, o, 1n+j, F32.add(acc, F32.mul(av, bv)), m, k, n, Array.get(F32, sa2, U32.from_nat(1n+j)), Array.get(F32, sb2, U32.from_nat(Nat.add(mm_bfloat(o, m, k, n), Nat.mul(1n+j, n)))), out_a) # --- fused backward kernels (single pass; F32 locals are reusable) --- def apply_bwd1(+op: Bwd1Op, +x: F32, +gy: F32) -> F32: match op: case BSigmoid{}: +y = F32.div(1.0, F32.add(1.0, F32.exp(F32.neg(x)))) F32.mul(gy, F32.mul(y, F32.sub(1.0, y))) case BTanh{}: +y = F32.tanh(x) F32.mul(gy, F32.sub(1.0, F32.mul(y, y))) case BSqrt{}: F32.div(gy, F32.mul(F32.sqrt(x), 2.0)) case BExp{}: F32.mul(gy, F32.exp(x)) case BNeg{}: F32.neg(gy) case BLog{}: F32.div(gy, x) case BAbs{}: Bool.pick(F32, F32.is_lt(x, 0.0), F32.neg(gy), gy) case BRelu{}: Bool.pick(F32, F32.is_gt(x, 0.0), gy, 0.0) case BId1{}: gy def ew1b(+op: Bwd1Op, rem: Nat, +idx: Nat, ra: Array & F32, rg: Array & F32, out_a: Array) -> Array: match rem: case 0n: out_a case 1n+p: (a2, x) = ra (g2, gy) = rg ew1b(op, p, 1n+idx, Array.get(F32, a2, U32.from_nat(1n+idx)), Array.get(F32, g2, U32.from_nat(1n+idx)), Array.set(F32, out_a, U32.from_nat(idx), apply_bwd1(op, x, gy))) def apply_bwd2(+op: Bwd2Op, +x: F32, +y: F32, +gy: F32) -> F32: match op: case BDivL{}: F32.div(gy, y) case BDivR{}: F32.neg(F32.mul(gy, F32.div(x, F32.mul(y, y)))) case BPowL{}: fwd = F32.pow(x, y) F32.div(F32.mul(F32.mul(gy, y), fwd), x) case BPowR{}: F32.mul(F32.mul(gy, F32.pow(x, y)), F32.log(x)) case BMax2L{}: Bool.pick(F32, F32.is_gt(x, y), gy, Bool.pick(F32, F32.is_eq(x, y), F32.div(gy, 2.0), 0.0)) case BMax2R{}: Bool.pick(F32, F32.is_lt(x, y), gy, Bool.pick(F32, F32.is_eq(x, y), F32.div(gy, 2.0), 0.0)) case BId{}: gy case B2Neg{}: F32.neg(gy) case BMulL{}: F32.mul(y, gy) case BMulR{}: F32.mul(x, gy) case BZero{}: 0.0 case BTriL{}: Bool.pick(F32, F32.is_gt(x, 0.5), gy, 0.0) case BTriR{}: Bool.pick(F32, F32.is_gt(x, 0.5), 0.0, gy) def ew2b(+op: Bwd2Op, rem: Nat, +idx: Nat, +sa: List<&2, Nat>, +sb: List<&2, Nat>, +out: List<&2, Nat>, ra: Array & F32, rb: Array & F32, rg: Array & F32, out_a: Array) -> Array: match rem: case 0n: out_a case 1n+p: (a2, x) = ra (b2, y) = rb (g2, gy) = rg ew2b(op, p, 1n+idx, sa, sb, out, ew2_next(a2, sa, out, idx), ew2_next(b2, sb, out, idx), ew2_next(g2, out, out, idx), Array.set(F32, out_a, U32.from_nat(idx), apply_bwd2(op, x, y, gy))) # ---------- eager-backward dispatchers used by autograd.bend ---------- def leaf_wrap(a: Array, +shape: List<&2, Nat>) -> Tensor: Leaf{a, shape} def sub_tag(+side: Nat) -> Bwd2Op: match side: case 0n: BId{} case _: B2Neg{} def mul_tag(+side: Nat) -> Bwd2Op: match side: case 0n: BMulL{} case _: BMulR{} def div_tag(+side: Nat) -> Bwd2Op: match side: case 0n: BDivL{} case _: BDivR{} def pow_tag(+side: Nat) -> Bwd2Op: match side: case 0n: BPowL{} case _: BPowR{} def max_tag(+side: Nat) -> Bwd2Op: match side: case 0n: BMax2L{} case _: BMax2R{} def un_tag(+op: UnOp) -> Bwd1Op: match op: case OpNeg{}: BNeg{} case OpExp{}: BExp{} case OpLog{}: BLog{} case OpSqrt{}: BSqrt{} case OpAbs{}: BAbs{} case OpTanh{}: BTanh{} case OpSigmoid{}: BSigmoid{} case OpRelu{}: BRelu{} def bin_tag(+op: BinOp, +side: Nat) -> Bwd2Op: match op: case OpAdd{}: BId{} case OpSub{}: sub_tag(side) case OpMul{}: mul_tag(side) case OpDiv{}: div_tag(side) case OpPow{}: pow_tag(side) case OpMax2{}: max_tag(side) case _: BZero{} def ew2b_gen(+op: BinOp, +side: Nat, rem: Nat, +idx: Nat, out_a: Array, ra: Array & F32, rb: Array & F32, rg: Array & F32) -> Array: ew2b(bin_tag(op, side), rem, idx, Nil{}, Nil{}, Nil{}, ra, rb, rg, out_a) def ew1b_gen(+op: UnOp, rem: Nat, +idx: Nat, out_a: Array, ra: Array & F32, rg: Array & F32) -> Array: ew1b(un_tag(op), rem, idx, ra, rg, out_a) # transpose of the last two axes (matmul VJP) def lead_of(+s: List<&2, Nat>, +keep: Nat) -> List<&2, Nat>: match s keep: case _, 0n: Nil{} case Nil{}, _: Nil{} case Con{h, t}, 1n+k2: h <> lead_of(t, k2) def tlast_order(+s: List<&2, Nat>) -> List<&2, Nat>: +r = List.length(&2, Nat, s) List.append(&2, Nat, lead_of(s, Nat.sub(r, 2n)), (Nat.sub(r, 1n) <> Nat.sub(r, 2n) <> Nil{} : List<&2, Nat>)) def tlast_shape(+s: List<&2, Nat>) -> List<&2, Nat>: permute_shape(s, tlast_order(s)) # inverse of a permutation (permute VJP) def insert_nat(+xs: List<&2, Nat>, +k: Nat, +v: Nat) -> List<&2, Nat>: match xs k: case _, 0n: v <> xs case Nil{}, _: v <> Nil{} case Con{h, t}, 1n+k2: h <> insert_nat(t, k2, v) def perm_inv_go(+order: List<&2, Nat>, +i: Nat) -> List<&2, Nat>: match order: case Nil{}: Nil{} case Con{h, t}: insert_nat(perm_inv_go(t, 1n+i), h, i) def perm_inv(+order: List<&2, Nat>) -> List<&2, Nat>: perm_inv_go(order, 0n) # ---------- the dispatcher: self-recursion only, kernels never call back ---------- def ev(+ps: List<&2, F32>, t: Tensor) -> Array: match t: case Leaf{buf, _}: buf case Param{+pid, +off, +shape}: params_get(ps, off, Sh.numel(shape)) case Un{+op, a, +out}: src = ev(ps, a) +n = Sh.numel(out) ew1(op, n, 0n, Array.get(F32, src, 0), Array.new(F32, depth_for(n), 0.0)) case Bin{+op, a, b, +sa, +sb, +out}: src_a = ev(ps, a) src_b = ev(ps, b) +total = Sh.numel(out) ew2(op, total, 0n, sa, sb, out, Array.get(F32, src_a, 0), Array.get(F32, src_b, 0), Array.new(F32, depth_for(total), 0.0)) case Tri{c, a, b, +sc, +sa, +sb, +out}: src_c = ev(ps, c) src_a = ev(ps, a) src_b = ev(ps, b) +total = Sh.numel(out) ew3(total, 0n, sc, sa, sb, out, Array.get(F32, src_c, 0), Array.get(F32, src_a, 0), Array.get(F32, src_b, 0), Array.new(F32, depth_for(total), 0.0)) case Mat{a, b, +batch, +m, +k, +n, +out}: src_a = ev(ps, a) src_b = ev(ps, b) +total = Sh.numel(out) mmloop(total, k, 0n, 0n, 0.0, m, k, n, Array.get(F32, src_a, U32.from_nat(mm_aflat(0n, m, k, n))), Array.get(F32, src_b, U32.from_nat(mm_bflat(0n, m, k, n))), Array.new(F32, depth_for(total), 0.0)) case Mov{+op, a, +arg1, +arg2, +sin, +out}: src = ev(ps, a) +total = Sh.numel(out) movloop(op, total, 0n, arg1, arg2, sin, out, Array.new(F32, depth_for(total), 0.0), Array.get(F32, src, U32.from_nat(mov_src_idx(op, 0n, arg1, arg2, sin, out)))) case Red{+op, a, +axis, +sin, +out}: src = ev(ps, a) +total = Sh.numel(out) +dim = Sh.at(sin, axis) stride = Sh.numel(stride_tail(sin, axis)) redloop(op, total, dim, 0n, 0n, red_base(0n, out, sin, axis), stride, dim, init_acc(op), axis, sin, out, Array.new(F32, depth_for(total), 0.0), Array.get(F32, src, U32.from_nat(red_base(0n, out, sin, axis)))) case Bwd1{+op, a, g, +sin, +out}: src_a = ev(ps, a) src_g = ev(ps, g) +total = Sh.numel(out) ew1b(op, total, 0n, Array.get(F32, src_a, 0), Array.get(F32, src_g, 0), Array.new(F32, depth_for(total), 0.0)) case Bwd2{+op, a, b, g, +sa, +sb, +out}: src_a = ev(ps, a) src_b = ev(ps, b) src_g = ev(ps, g) +total = Sh.numel(out) ew2b(op, total, 0n, sa, sb, out, Array.get(F32, src_a, 0), Array.get(F32, src_b, 0), Array.get(F32, src_g, 0), Array.new(F32, depth_for(total), 0.0)) # ---------- printing helper for toys ---------- def show_flat(rem: Nat, +idx: Nat, r: Array & F32, +acc: String) -> String: match rem: case 0n: acc case 1n+p: (a2, x) = r show_flat(p, 1n+idx, Array.get(F32, a2, U32.from_nat(1n+idx)), String.append(String.append(acc, F32.show(x)), " "))