LAWS.bend source
LAWS.bend on the hub · documented module
# 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 Baseimport ./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}