# tinygrad/shape.bend -- pure shape arithmetic for tinygrad-bend. # # Port of the shape contracts of upstream tinygrad: # - broadcasting: tensor.py `_broadcasted` / UOp broadcast (align right, 1s expand) # - matmul: mixin/op.py `dot` ((...,m,k) @ (...,k,n) -> (...,m,n)) # - reduce: mixin/reduce.py `_reduce` (drop the reduced dim) # - pad/shrink: mixin/movement.py (pad adds, shrink removes; grad duality) # # Everything in this file is total, terminating, and free of F32, so LAWS.bend # can state laws about it and PROOF.bend can prove them. Dims are >= 1: empty # (size-0) tensors are not modeled (documented deviation, see README). import Base # Shapes are lists of dims, row-major (C order), like tinygrad. # def Shape: List<&2, Nat> (kept inline; List below means List<&2, Nat>) # numel -- product of dims. Mirrors UOp.shape prod / Tensor numel. def numel(+s: List<&2, Nat>) -> Nat: match s: case Nil{}: 1n case Con{h, t}: Nat.mul(h, numel(t)) # maxn -- structural max of two Nats. Own def (not Base's Nat.max, which is # pick/cmp-based and resists induction); laws talk about this one. def maxn(+a: Nat, +b: Nat) -> Nat: match a b: case 0n, _: b case 1n+a2, 0n: 1n+a2 case 1n+a2, 1n+b2: 1n+maxn(a2, b2) # bcast -- common broadcast shape of s and t (numpy-style: ranks align right, # each dim is the max; a dim of 1 stretches). Mirrors tinygrad `_broadcasted`. def bcast(+s: List<&2, Nat>, +t: List<&2, Nat>) -> List<&2, Nat>: match s t: case Nil{}, Nil{}: Nil{} case Nil{}, Con{b, bs}: b <> bs case Con{a, as}, Nil{}: a <> as case Con{a, as}, Con{b, bs}: maxn(a, b) <> bcast(as, bs) # bcast_ok -- can s and t broadcast against each other? def dim_ok(c: Cmp, d: Nat, e: Nat) -> Bool: match c: case LT{}: Nat.is_eq(e, 1n) case EQ{}: True{} case GT{}: Nat.is_eq(d, 1n) def bcast_ok(+s: List<&2, Nat>, +t: List<&2, Nat>) -> Bool: match s t: case Nil{}, _: True{} case _, Nil{}: True{} case Con{a, as}, Con{b, bs}: Bool.and(dim_ok(Nat.cmp(a, b), a, b), bcast_ok(as, bs)) # bdim2 -- the dominance core: d is the peeled dim. d2 == 0n means the source # dim was 1 (stretches to anything); otherwise the dims must shrink together. def bdim2(d2: Nat, e2: Nat) -> Type: match d2 e2: case 0n, _: Unit case 1n+d3, 0n: Empty case 1n+d3, 1n+e3: bdim2(d3, e3) def bdim(d: Nat, e: Nat) -> Type: match d e: case 0n, _: Empty case 1n+d2, 0n: Empty case 1n+d2, 1n+e2: bdim2(d2, e2) def bc(+s: List<&2, Nat>, +t: List<&2, Nat>) -> Type: match s t: case Nil{}, _: Unit case Con{_, _}, Nil{}: {numel(s) == 1n : Nat} case Con{a, as}, Con{b, bs}: bdim(a, b) & bc(as, bs) # matmul_out -- shape law of dot: (m,k) @ (k,n) -> (m,n). def matmul_out(+m: Nat, +k: Nat, +n: Nat) -> List<&2, Nat>: m <> n <> Nil{} # pad_shape -- pad each dim: new = lo + dim + hi (arg = los, his lists). def pad_shape(+s: List<&2, Nat>, +los: List<&2, Nat>, +his: List<&2, Nat>) -> List<&2, Nat>: match s los his: case Nil{}, _, _: Nil{} case _, Nil{}, _: s case _, _, Nil{}: s case Con{h, t}, Con{lo, los2}, Con{hi, his2}: Nat.add(lo, Nat.add(h, hi)) <> pad_shape(t, los2, his2) # shrink_shape -- inverse view: new = dim - lo - hi. def shrink_shape(+s: List<&2, Nat>, +los: List<&2, Nat>, +his: List<&2, Nat>) -> List<&2, Nat>: match s los his: case Nil{}, _, _: Nil{} case _, Nil{}, _: s case _, _, Nil{}: s case Con{h, t}, Con{lo, los2}, Con{hi, his2}: Nat.sub(Nat.sub(h, lo), hi) <> shrink_shape(t, los2, his2) # drop_at -- shape of a reduction over axis k (0-based): removes dim k. # Total: k >= rank leaves the shape untouched (callers never do that). def drop_at(+s: List<&2, Nat>, +k: Nat) -> List<&2, Nat>: match s k: case Nil{}, _: Nil{} case Con{h, t}, 0n: t case Con{h, t}, 1n+k2: h <> drop_at(t, k2) # at -- dim k of s (1 when out of range), so reduce_numel below is total. def at(+s: List<&2, Nat>, +k: Nat) -> Nat: match s k: case Nil{}, _: 1n case Con{h, t}, 0n: h case Con{h, t}, 1n+k2: at(t, k2) def bdim_idx_if(+d: Nat, +i: Nat, z: Bool) -> Nat: match z: case True{}: 0n case False{}: i # flat -- row-major flat index of a multi-index (multi-index lists are short; # rank is bounded by nesting depth of user code, so O(rank) walks are fine). def flat(+mi: List<&2, Nat>, +s: List<&2, Nat>) -> Nat: match mi s: case Nil{}, _: 0n case _, Nil{}: 0n case Con{i, it}, Con{d, ds}: Nat.add(Nat.mul(i, numel(ds)), flat(it, ds)) def bdim_idx(+d: Nat, +i: Nat) -> Nat: bdim_idx_if(d, i, Nat.is_eq(d, 1n)) # bdim_idx -- index into a broadcast dim: a stretched dim (size 1) reads 0. # bcast_flat -- the source flat index that output multi-index `mi` reads from, # when out shape is `os` and source shape is `s` (same rank; s dims may be 1). def bcast_flat_skip(+ods: List<&2, Nat>) -> Nat: match ods: case Nil{}: 0n case Con{_, t}: bcast_flat_skip(t) def bcast_flat_go(+mi: List<&2, Nat>, +s: List<&2, Nat>, +os: List<&2, Nat>) -> Nat: match mi s os: case Nil{}, _, _: 0n case _, Nil{}, Nil{}: 0n case _, Nil{}, Con{_, ods}: # right-aligned rank difference: leading output dims read past the # (implicitly 1-sized) source dims -- nothing to add, keep walking bcast_flat_skip(ods) case _, _, Nil{}: 0n case Con{i, it}, Con{d, ds}, Con{od, ods}: Nat.add(Nat.mul(bdim_idx(d, i), numel(ds)), bcast_flat_go(it, ds, ods)) def bcast_flat(+mi: List<&2, Nat>, +s: List<&2, Nat>, +os: List<&2, Nat>) -> Nat: bcast_flat_go(mi, s, os) # unflat -- multi-index of flat index i in shape s. Head index = i / numel(rest), # tail recurses on the remainder (row-major; wraps into range, so total). def unflat_go(+ds: List<&2, Nat>, +i: Nat) -> List<&2, Nat>: match ds: case Nil{}: Nil{} case Con{d2, ds2}: Nat.div(i, numel(ds2)) <> unflat_go(ds2, Nat.mod(i, numel(ds2))) def unflat(+i: Nat, +s: List<&2, Nat>) -> List<&2, Nat>: unflat_go(s, i)