# LAWS.bend -- the laws of tinygrad, stated for the Bend port. # # The human states these; PROOF.bend must prove them. Each law cites the # upstream tinygrad code it captures (see LAW.md for the full mapping). # Only F32-free, structural facts are stated here: Bend's F32 is axiomatic, # so numeric correctness is pinned by the runtime toys in tests/ instead. import Base import ./tinygrad/shape.bend as Sh # LAW 1 (broadcast commutes) -- tinygrad tensor.py `_broadcasted`: both # operands are broadcast to ONE common shape; which operand came first # cannot change it. law bcast_commute: for +s: List<&2, Nat> for +t: List<&2, Nat> {Sh.bcast(s, t) == Sh.bcast(t, s) : List<&2, Nat>} # LAW 2 (matmul shape) -- tinygrad mixin/op.py `dot`: (m,k) @ (k,n) = (m,n); # the inner dims must agree, the output carries the outer dims. law matmul_shape: for +m: Nat for +k: Nat for +n: Nat {Sh.matmul_out(m, k, n) == m <> n <> Nil{} : List<&2, Nat>} # LAW 3 (broadcast numel) -- the shape half of tinygrad's broadcast gradient # rule (mixin/gradient.py: a shaped edge's gradient is summed back to its # source's shape): when t dominates s, broadcasting s against t yields # exactly t's element count, so summing a t-shaped gradient over the # broadcast axes can restore s's shape. law bcast_numel: for +s: List<&2, Nat> for +t: List<&2, Nat> for c: Sh.bc(s, t) {Sh.numel(Sh.bcast(s, t)) == Sh.numel(t) : Nat} # LAW 4 (gradient accumulation is lossless) -- tinygrad tensor.py `backward` # (`t.grad.assign(t.grad + g)`): when a tensor is used several times, every # use contributes its gradient and none is dropped; the port's gradient table # accumulates a contribution list per parameter, so the invariant is that # appending contributions preserves their count. law grad_accum_len: for +xs: List<&2, Nat> for +ys: List<&2, Nat> {List.length(&2, Nat, List.append(&2, Nat, xs, ys)) == Nat.add(List.length(&2, Nat, xs), List.length(&2, Nat, ys)) : Nat} # LAW 6 (reduce count, axis 0) -- tinygrad mixin/reduce.py `_reduce`: summing # over axis 0 removes dim 0, so numel(s) = dim_0(s) * numel(reduced). The # general-axis form needs mul-associativity, which the port pins at runtime # in tests/toy_reduce.bend instead (see LAW.md). law reduce_numel: for +s: List<&2, Nat> {Sh.numel(s) == Nat.mul(Sh.at(s, 0n), Sh.numel(Sh.drop_at(s, 0n))) : Nat}