From b17e15d1aacd03cda65654eb8869fcc4988cb296 Mon Sep 17 00:00:00 2001 From: George Hotz Date: Wed, 11 Mar 2026 15:22:46 +0800 Subject: [PATCH] support ranges on call --- extra/callrange/test.py | 18 ++++++++++++++++++ tinygrad/uop/ops.py | 16 +++++++++++----- 2 files changed, 29 insertions(+), 5 deletions(-) create mode 100644 extra/callrange/test.py diff --git a/extra/callrange/test.py b/extra/callrange/test.py new file mode 100644 index 0000000000..0454d2a04c --- /dev/null +++ b/extra/callrange/test.py @@ -0,0 +1,18 @@ +from tinygrad import UOp, dtypes, Device, Tensor + +if __name__ == "__main__": + B0 = UOp.new_buffer(Device.DEFAULT, 100, dtypes.float).reshape(10,10) + B1 = UOp.new_buffer(Device.DEFAULT, 100, dtypes.float).reshape(10,10) + + R0 = UOp.range(10, axis_id=0) + R1 = UOp.range(10, axis_id=1) + + b0 = UOp.param(0, dtypes.float, (10,10)) + b1 = UOp.param(1, dtypes.float, (10,10)) + r0 = UOp.param(2, dtypes.index, ()) + r1 = UOp.param(3, dtypes.index, ()) + + fxn = (b0[r0, r1] + b1[r0, r1]).call(B0, B1, R0, R1) + t = Tensor(fxn) + t.realize() + diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index 23230d0abb..91bc70d687 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -207,7 +207,7 @@ class UOp(OpMixin, metaclass=UOpMetaClass): def _shape(self) -> tuple[sint, ...]|None: match self.op: # late ops don't have shape - case Ops.UNIQUE | Ops.LUNIQUE | Ops.DEVICE | Ops.RANGE | Ops.LOAD | Ops.IF | Ops.BARRIER | Ops.CUSTOM | Ops.CUSTOMI | \ + case Ops.UNIQUE | Ops.LUNIQUE | Ops.DEVICE | Ops.LOAD | Ops.IF | Ops.BARRIER | Ops.CUSTOM | Ops.CUSTOMI | \ Ops.VECTORIZE | Ops.GEP | Ops.SPECIAL | Ops.UNROLL | Ops.CONTRACT | Ops.SINK | \ Ops.LINEAR | Ops.PROGRAM | Ops.SOURCE | Ops.BINARY | Ops.INS: return None @@ -218,15 +218,17 @@ class UOp(OpMixin, metaclass=UOpMetaClass): return None case Ops.INDEX: - # non pointer index doesn't have a shape - if not isinstance(self.dtype, PtrDType): return None + # non pointer index + if not isinstance(self.dtype, PtrDType): + idxs = flatten([d.shape for d in self.src[1:]]) + return tuple(idxs) + self.src[0].shape[len(self.src)-1:] # fully indexed doesn't have a shape. TODO: remove this if self.src[0]._shape is None or len(self.src[1:]) == len(self.src[0].shape): return None # pointer index return self.src[0].shape[len(self.src[1:]):] # some ops init the shape - case Ops.CONST | Ops.VCONST | Ops.DEFINE_VAR | Ops.BIND: return () + case Ops.CONST | Ops.VCONST | Ops.DEFINE_VAR | Ops.BIND | Ops.RANGE: return () case Ops.BUFFER: return (self.arg,) case Ops.BUFFER_VIEW: return (self.arg[0],) case Ops.CUSTOM_FUNCTION: return None @@ -246,7 +248,10 @@ class UOp(OpMixin, metaclass=UOpMetaClass): inner_shape = self.src[0]._shape if inner_shape is None: return None # substitute internal PARAMs in the shape with corresponding args - return tuple(graph_rewrite(s, _pm_resolve_params, self.src[1:], walk=True) if isinstance(s, UOp) else s for s in inner_shape) + ret = tuple(graph_rewrite(s, _pm_resolve_params, self.src[1:], walk=True) if isinstance(s, UOp) else s for s in inner_shape) + # NOTE: this requires the RANGEs directly on the call + prepend = tuple([x.vmax+1 for x in self.src[1:] if x.op is Ops.RANGE]) + return prepend+ret # TODO: disallow shape changing bitcast case Ops.BITCAST: @@ -1462,6 +1467,7 @@ renderer = PatternMatcher([ (UPat(Ops.PARAM, src=(UPat(), UPat(), UPat(), UPat(), UPat(Ops.NOOP, name="x"))), lambda x: x.arg), (UPat((Ops.SPECIAL), name="x"), lambda x: x.arg), (UPat(Ops.RANGE, name="x"), lambda x: f"r{range_str(x)}"), + (UPat(Ops.PARAM, name="x"), lambda x: f"p{x.arg}"), (UPat((Ops.CONST, Ops.VCONST), name="x"), lambda x: str(x.arg)), (UPat(Ops.UNROLL, name="x"), lambda ctx,x,u: f"UNROLL({ctx[x.src[0]]}, {u.arg})"), (UPat(Ops.CAST, name="x"), lambda ctx,x: f"({str(x.dtype)[7:]})({ctx[x.src[0]]})"),