tinygrad/shape.bend checks
raw source on the hub · import 0xe6b82fa6c4c459c7023adf4b12a3ac46/tinygrad/shape.bend as Shape
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).
1 import
import Base
Definitions
def numel source · line 19 · raw
@+s:List<&2, Nat> -> Nat
numel -- product of dims. Mirrors UOp.shape prod / Tensor numel.
def maxn source · line 28 · raw
@+a:Nat -> @+b:Nat -> Nat
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 bcast source · line 39 · raw
@+s:List<&2, Nat> -> @+t:List<&2, Nat> -> List<&2, Nat>
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 dim_ok source · line 51 · raw
@c:Cmp -> @d:Nat -> @e:Nat -> Bool
bcast_ok -- can s and t broadcast against each other?
def bcast_ok source · line 60 · raw
@+s:List<&2, Nat> -> @+t:List<&2, Nat> -> Bool
def bdim2 source · line 71 · raw
@d2:Nat -> @e2:Nat -> Type
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 bdim source · line 80 · raw
@d:Nat -> @e:Nat -> Type
def bc source · line 89 · raw
@+s:List<&2, Nat> -> @+t:List<&2, Nat> -> Type
def matmul_out source · line 99 · raw
@+m:Nat -> @+k:Nat -> @+n:Nat -> List<&2, Nat>
matmul_out -- shape law of dot: (m,k) @ (k,n) -> (m,n).
def pad_shape source · line 103 · raw
@+s:List<&2, Nat> -> @+los:List<&2, Nat> -> @+his:List<&2, Nat> -> List<&2, Nat>
pad_shape -- pad each dim: new = lo + dim + hi (arg = los, his lists).
def shrink_shape source · line 115 · raw
@+s:List<&2, Nat> -> @+los:List<&2, Nat> -> @+his:List<&2, Nat> -> List<&2, Nat>
shrink_shape -- inverse view: new = dim - lo - hi.
def drop_at source · line 128 · raw
@+s:List<&2, Nat> -> @+k:Nat -> List<&2, Nat>
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 at source · line 138 · raw
@+s:List<&2, Nat> -> @+k:Nat -> Nat
at -- dim k of s (1 when out of range), so reduce_numel below is total.
def bdim_idx_if source · line 147 · raw
@+d:Nat -> @+i:Nat -> @z:Bool -> Nat
def flat source · line 156 · raw
@+mi:List<&2, Nat> -> @+s:List<&2, Nat> -> Nat
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 bdim_idx source · line 165 · raw
@+d:Nat -> @+i:Nat -> Nat
def bcast_flat_skip source · line 171 · raw
@+ods:List<&2, Nat> -> Nat
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_go source · line 178 · raw
@+mi:List<&2, Nat> -> @+s:List<&2, Nat> -> @+os:List<&2, Nat> -> Nat
def bcast_flat source · line 193 · raw
@+mi:List<&2, Nat> -> @+s:List<&2, Nat> -> @+os:List<&2, Nat> -> Nat
def unflat_go source · line 198 · raw
@+ds:List<&2, Nat> -> @+i:Nat -> List<&2, Nat>
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 source · line 204 · raw
@+i:Nat -> @+s:List<&2, Nat> -> List<&2, Nat>