~/bend-docscommunity

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>