From 0411b09763a105b7d2ee592a24bd61776e3940bc Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Thu, 8 May 2025 07:04:27 -0700 Subject: [PATCH] small changes from new multi [pr] (#10213) --- tinygrad/engine/grouper.py | 6 +++--- tinygrad/ops.py | 8 ++++---- tinygrad/viz/serve.py | 5 +++-- 3 files changed, 10 insertions(+), 9 deletions(-) diff --git a/tinygrad/engine/grouper.py b/tinygrad/engine/grouper.py index baa4a8dba4..2b63e68165 100644 --- a/tinygrad/engine/grouper.py +++ b/tinygrad/engine/grouper.py @@ -82,9 +82,9 @@ sym = symbolic_simple+PatternMatcher([ # split_reduceop (UPat(Ops.REDUCE_AXIS, name="reduce", src=(UPat.var("x"),)), split_reduceop), # COPY(CONST) creates a new CONST on the destination device - (UPat(Ops.COPY, name="root", src=(UPat.cvar("x"), UPat())), lambda root,x: root.const_like(x.arg)), + (UPat(Ops.COPY, name="root", src=(UPat.cvar("x"),), allow_any_len=True), lambda root,x: root.const_like(x.arg)), # store a shrink before COPY, otherwise view after the COPY - (UPat(Ops.COPY, src=(UPat(Ops.VIEW, name="v"), UPat()), name="copy"), lambda copy,v: v.contiguous().copy_to_device(copy.device) \ + (UPat(Ops.COPY, src=(UPat(Ops.VIEW, name="v"),), name="copy", allow_any_len=True), lambda copy,v: v.contiguous().copy_to_device(copy.device) \ if prod(v.shape) < prod(v.base.shape) else v.base.copy_to_device(copy.device).view(v.st)), # remove cast to image when it's already a contiguous image (UPat(Ops.CAST, name="cast", src=(UPat(Ops.VIEW, name="vm", src=(UPat(Ops.CONTIGUOUS, name="base"),)),)), @@ -143,7 +143,7 @@ do_realize = PatternMatcher([ # realize before expand or unsafe pad ops (UPat(Ops.VIEW, src=(UPat(GroupOp.All-ALWAYS_CONTIGUOUS, name="tr"),), name="view"), realize_before_view), # realize before COPY - (UPat(Ops.COPY, src=(UPat(GroupOp.All-ALWAYS_CONTIGUOUS, name="tr"), UPat())), realize), + (UPat(Ops.COPY, src=(UPat(GroupOp.All-ALWAYS_CONTIGUOUS, name="tr"),), allow_any_len=True), realize), ]) def recursive_group(tr:UOp, st:ShapeTracker, r:UOp, children:defaultdict[UOp, dict[UOp, None]], realizes:dict[UOp, None], diff --git a/tinygrad/ops.py b/tinygrad/ops.py index b515c760ba..399862ad64 100644 --- a/tinygrad/ops.py +++ b/tinygrad/ops.py @@ -490,9 +490,9 @@ class UOp(MathTrait, metaclass=UOpMetaClass): assert op is Ops.BIND, f"unknown op {op}" var, val = arg.unbind() return var.replace(src=(UOp(Ops.VIEW, dtypes.void, (UOp(Ops.DEVICE, arg=device),), ShapeTracker.from_shape(shape)),)).bind(val) - def copy_to_device(self, device:str|tuple[str, ...]): - if isinstance(device, tuple): return UOp(Ops.COPY, self.dtype, (self,)+tuple(UOp(Ops.DEVICE, arg=d) for d in device)) - return UOp(Ops.COPY, self.dtype, (self, UOp(Ops.DEVICE, arg=device))) + def copy_to_device(self, device:str|tuple[str, ...], arg=None): + if isinstance(device, tuple): return UOp(Ops.COPY, self.dtype, (self,)+tuple(UOp(Ops.DEVICE, arg=d) for d in device), arg) + return UOp(Ops.COPY, self.dtype, (self, UOp(Ops.DEVICE, arg=device)), arg) def clone(self) -> UOp: return self.copy_to_device(self.device) @property def metadata(self) -> tuple[Metadata, ...]|Metadata|None: return self.arg.metadata if self.op is Ops.KERNEL else all_metadata.get(self, None) @@ -536,7 +536,7 @@ class UOp(MathTrait, metaclass=UOpMetaClass): def _device(self) -> Optional[str|tuple[str, ...]]: if self.op is Ops.DEVICE: return self.arg if self.op is Ops.MULTI: return tuple(cast(str, x.device) for x in self.src) - if self.op is Ops.COPY: + if self.op in {Ops.COPY, Ops.BUFFER}: if len(self.src) > 2: return tuple(cast(str, x.device) for x in self.src[1:]) return self.src[1].device return dsrcs[0]._device if len(dsrcs:=[x for x in self.src if x._device is not None]) != 0 else None diff --git a/tinygrad/viz/serve.py b/tinygrad/viz/serve.py index 62b0a2fb20..2c1fccb37e 100755 --- a/tinygrad/viz/serve.py +++ b/tinygrad/viz/serve.py @@ -54,6 +54,7 @@ class GraphRewriteDetails(TypedDict): upat: tuple[tuple[str, int], str]|None # [loc, source_code] of the matched UPat def shape_to_str(s:tuple[sint, ...]): return "(" + ','.join(srender(x) for x in s) + ")" +def mask_to_str(s:tuple[tuple[sint, sint], ...]): return "(" + ','.join(shape_to_str(x) for x in s) + ")" def uop_to_json(x:UOp) -> dict[int, dict]: assert isinstance(x, UOp) @@ -69,8 +70,8 @@ def uop_to_json(x:UOp) -> dict[int, dict]: if u in excluded: continue argst = str(u.arg) if u.op is Ops.VIEW: - argst = ("\n".join([f"{shape_to_str(v.shape)} / {shape_to_str(v.strides)}"+(f"\nMASK {v.mask}" if v.mask is not None else "")+ - ("" if v.offset == 0 else f" / {srender(v.offset)}") for v in unwrap(u.st).views])) + argst = ("\n".join([f"{shape_to_str(v.shape)} / {shape_to_str(v.strides)}"+("" if v.offset == 0 else f" / {srender(v.offset)}")+ + (f"\nMASK {mask_to_str(v.mask)}" if v.mask is not None else "") for v in unwrap(u.st).views])) label = f"{str(u.op).split('.')[1]}{(chr(10)+word_wrap(argst.replace(':', ''))) if u.arg is not None else ''}" if u.dtype != dtypes.void: label += f"\n{u.dtype}" for idx,x in enumerate(u.src):