# PROOF.bend -- proofs of the laws in LAWS.bend. The gate: # bend PROOF.bend => "All terms check." # # Proof notes (learned on Bend 2.0.4): # - induction is a recursive call on a structurally smaller pattern variable # - `%e : P` rewrites the goal: P is the goal with `_` marking the b-side of # e's equation, and the goal becomes P with `a` in that place # - evidence of `Empty` (clashes like 2n-vs-3n broadcast) is discharged by a # match with no cases import Base import ./tinygrad/shape.bend as Sh import ./LAWS.bend as Laws # ---------- LAW 1: broadcast commutes ---------- def max_comm(+a: Nat, +b: Nat) -> {Sh.maxn(a, b) == Sh.maxn(b, a) : Nat}: match a b: case 0n, 0n: {==} case 0n, 1n+b2: {==} case 1n+a2, 0n: {==} case 1n+a2, 1n+b2: %max_comm(a2, b2) : {1n+Sh.maxn(a2, b2) == 1n+_ : Nat} {==} def Laws.bcast_commute(s, t): match s t: case Nil{}, Nil{}: {==} case Nil{}, Con{b, bs}: {==} case Con{a, as}, Nil{}: {==} case Con{a, as}, Con{b, bs}: %max_comm(a, b) : {Sh.maxn(a, b) <> Sh.bcast(as, bs) == _ <> Sh.bcast(bs, as) : List<&2, Nat>} %Laws.bcast_commute(as, bs) : {Sh.maxn(a, b) <> Sh.bcast(as, bs) == Sh.maxn(a, b) <> _ : List<&2, Nat>} {==} # ---------- LAW 2: matmul shape ---------- def Laws.matmul_shape(m, k, n): {==} # ---------- LAW 3: broadcast numel ---------- # a dominating dim yields the max def max_bdim2(d2: Nat, e2: Nat, ev: Sh.bdim2(d2, e2)) -> {Sh.maxn(d2, e2) == e2 : Nat}: match d2 e2: case 0n, 0n: {==} case 0n, 1n+e3: {==} case 1n+d3, 0n: match ev: case 1n+d3, 1n+e3: %max_bdim2(d3, e3, ev) : {1n+Sh.maxn(d3, e3) == 1n+_ : Nat} {==} def max_bdim(d: Nat, e: Nat, ev: Sh.bdim(d, e)) -> {Sh.maxn(d, e) == e : Nat}: match d e: case 0n, _: match ev: case 1n+d2, 0n: match ev: case 1n+d2, 1n+e2: Equal.cong(Nat, Nat, x => 1n+x, Sh.maxn(d2, e2), e2, max_bdim2(d2, e2, ev)) def Laws.bcast_numel(s, t, c): match s t: case Nil{}, Nil{}: {==} case Nil{}, Con{b, bs}: {==} case Con{_, _}, Nil{}: c case Con{a, as}, Con{b, bs}: (dev, rest) = c ih = Laws.bcast_numel(as, bs, rest) c1 = Equal.cong(Nat, Nat, x => (Nat.mul(Sh.maxn(a, b), x) : Nat), Sh.numel(Sh.bcast(as, bs)), Sh.numel(bs), ih) c2 = Equal.cong(Nat, Nat, x => (Nat.mul(x, Sh.numel(bs)) : Nat), Sh.maxn(a, b), b, max_bdim(a, b, dev)) Equal.trans(Nat, Nat.mul(Sh.maxn(a, b), Sh.numel(Sh.bcast(as, bs))), Nat.mul(Sh.maxn(a, b), Sh.numel(bs)), Nat.mul(b, Sh.numel(bs)), c1, c2) # ---------- LAW 4: gradient accumulation is lossless ---------- def Laws.grad_accum_len(xs, ys): match xs: case Nil{}: {==} case Con{h, t}: Equal.cong(Nat, Nat, x => 1n+x, List.length(&2, Nat, List.append(&2, Nat, t, ys)), Nat.add(List.length(&2, Nat, t), List.length(&2, Nat, ys)), Laws.grad_accum_len(t, ys)) def sub_diag(+a: Nat) -> {Nat.sub(a, a) == 0n : Nat}: match a: case 0n: {==} case 1n+p: %sub_diag(p) : {Nat.sub(p, p) == _ : Nat} {==} def add_zero_right(+a: Nat) -> {Nat.add(a, 0n) == a : Nat}: match a: case 0n: {==} case 1n+p: %add_zero_right(p) : {1n+Nat.add(p, 0n) == 1n+_ : Nat} {==} # ---------- LAW 6: reduce count (axis 0) ---------- def Laws.reduce_numel(s): match s: case Nil{}: {==} case Con{h, t}: {==}