diff --git a/test/null/test_pattern_matcher.py b/test/null/test_pattern_matcher.py index 4e6819e7cb..b46d1862e1 100644 --- a/test/null/test_pattern_matcher.py +++ b/test/null/test_pattern_matcher.py @@ -69,7 +69,7 @@ class TestPatternMatcher(unittest.TestCase): def test_uop_set(self): matcher = PatternMatcher([(UPat((Ops.CONST, Ops.CAST), name="x"), lambda x: x.rtag())]) c1 = UOp.const(dtypes.bool, False) - c2 = UOp(Ops.CAST, dtypes.int, (c1,)) + c2 = UOp(Ops.CAST, arg=dtypes.int, src=(c1,)) c3 = UOp.const(dtypes.float, 1.0) c4 = UOp(Ops.ADD, src=(c3, c3)) self.assertEqual(matcher.rewrite(c1), c1.rtag()) diff --git a/tinygrad/codegen/decomp/dtype.py b/tinygrad/codegen/decomp/dtype.py index 15169256f6..18a7404f9c 100644 --- a/tinygrad/codegen/decomp/dtype.py +++ b/tinygrad/codegen/decomp/dtype.py @@ -1,7 +1,8 @@ +from dataclasses import replace from tinygrad.dtype import dtypes, DType, truncate from tinygrad.helpers import flatten, DEBUG, EMULATED_DTYPES, Context, SPEC from tinygrad.uop import GroupOp -from tinygrad.uop.ops import UOp, UPat, Ops, PatternMatcher, graph_rewrite +from tinygrad.uop.ops import UOp, UPat, Ops, PatternMatcher, graph_rewrite, ParamArg from tinygrad.renderer import Renderer from tinygrad.codegen.decomp.transcendental import exponent_bias, shl, shr @@ -120,7 +121,7 @@ def f2f_store(st, idx, val, fr:DType, to:DType): pm_long_decomp = PatternMatcher([ (UPat(GroupOp.Defines, src=(UPat.var("sz"),), name="x"), lambda x,sz: - x.replace(dtype=l2i_dt[x.dtype], src=(sz*2,)) if x.dtype in l2i_dt else None), + x.replace(dtype=l2i_dt[x.dtype], arg=replace(x.arg, dtype=l2i_dt[x.dtype]), src=(sz*2,)) if x.dtype in l2i_dt else None), (UPat(Ops.INDEX, tuple(l2i_dt.keys()), name='x'), lambda x: reindex(x, x.tag).replace(dtype=l2i_dt[x.dtype]) if x.tag is not None else None), (UPat(Ops.STORE, src=(UPat.var('idx'), UPat.var('val', tuple(l2i_dt.keys()))), name='st'), lambda st,idx,val: st.replace(src=(idx.rtag(0), val.rtag(0))).group(st.replace(src=(idx.rtag(1), val.rtag(1)))) if val.tag is None else None), @@ -144,7 +145,7 @@ pm_long_decomp = PatternMatcher([ # float decomposition patterns - ctx is (fr, to) tuple pm_float_decomp = PatternMatcher([ (UPat((*GroupOp.Defines, Ops.INDEX, Ops.SHRINK), name="x"), lambda ctx,x: - x.replace(dtype=f2f_dt[ctx[0]], tag=ctx[0]) + x.replace(dtype=f2f_dt[ctx[0]], arg=replace(x.arg, dtype=f2f_dt[ctx[0]]) if isinstance(x.arg, ParamArg) else x.arg, tag=ctx[0]) if x.dtype == ctx[0] and (x.op is not Ops.INDEX or x.src[0].op not in {Ops.LOAD, Ops.STACK}) else None), (UPat(Ops.LOAD, dtypes.floats, name="x"), lambda ctx,x: f2f_load(x, *ctx) if x.dtype == ctx[0] else None), # bitcasted load should just replace load diff --git a/tinygrad/renderer/isa/x86.py b/tinygrad/renderer/isa/x86.py index fc9bcc9708..1f0d0f1fc8 100644 --- a/tinygrad/renderer/isa/x86.py +++ b/tinygrad/renderer/isa/x86.py @@ -195,19 +195,19 @@ pre_isel_matcher = PatternMatcher([ (UPat(Ops.SHRINK, src=(UPat(), UPat(), UPat.cvar("c"))).load(allow_any_len=True, name="x"), lambda x,c: x.replace(dtype=x.dtype.scalar().vec(c.arg)) if c.arg > x.dtype.count else None), (UPat(Ops.STACK, name="x"), lambda x: x.replace(dtype=x.dtype.scalar().vec(len(x.src))) if 1 < len(x.src) != x.dtype.count else None), - (UPat(GroupOp.ALU.union({Ops.CAST, Ops.BITCAST}), name="x"), lambda x: x.replace(dtype=x.dtype.scalar().vec(c)) \ + (UPat(GroupOp.ALU.union({Ops.CAST, Ops.BITCAST}), name="x"), lambda x: x.replace(arg=x.dtype.scalar().vec(c)) \ if (c:=max([s.dtype.count for s in x.src], default=1)) > x.dtype.count else None), # zero extending scalar 32bit int is a noop - (UPat.var("y", dtypes.uint32).cast(dtypes.int64s, name="x"), lambda y,x: x.replace(op=Ops.NOOP) if y.dtype.count == 1 else None), + (UPat.var("y", dtypes.uint32).cast(dtypes.int64s, name="x"), lambda y,x: x.replace(op=Ops.NOOP, arg=None) if y.dtype.count == 1 else None), # cast between signed and unsigned int is a noop (UPat.var("y", dtypes.ints+(dtypes.bool,)).cast(dtypes.ints, name="x"), - lambda y,x: x.replace(op=Ops.NOOP) if x.dtype.itemsize == y.dtype.itemsize else None), + lambda y,x: x.replace(op=Ops.NOOP, arg=None) if x.dtype.itemsize == y.dtype.itemsize else None), # cast to < scalar int is a noop (UPat.var("y", dtypes.ints).cast(dtypes.ints, name="x"), - lambda y,x: x.replace(op=Ops.NOOP) if x.dtype.itemsize < y.dtype.itemsize and y.dtype.count == 1 else None), + lambda y,x: x.replace(op=Ops.NOOP, arg=None) if x.dtype.itemsize < y.dtype.itemsize and y.dtype.count == 1 else None), # bitcasts between scalar floats and ints are real, rest are noops (UPat.var("y").bitcast().named("x"), lambda y,x: None if y.dtype in dtypes.floats and x.dtype in dtypes.ints or \ - y.dtype in dtypes.ints and x.dtype in dtypes.floats else x.replace(op=Ops.NOOP)), + y.dtype in dtypes.ints and x.dtype in dtypes.floats else x.replace(op=Ops.NOOP, arg=None)), # noop of a noop is removed (UPat(Ops.NOOP, src=(UPat(Ops.NOOP),), name="x"), lambda x: x.replace(src=x.src[0].src)), # moving elements of a single register to another without shuffling is a noop diff --git a/tinygrad/schedule/rangeify.py b/tinygrad/schedule/rangeify.py index 6a8db1b4ec..bf4a5e6e49 100644 --- a/tinygrad/schedule/rangeify.py +++ b/tinygrad/schedule/rangeify.py @@ -377,7 +377,7 @@ def bufferize_to_store(ctx:itertools.count, x:UOp, idx:UOp, allow_locals=True): # NOTE: the local BUFFER needs to be disambiguated here if x.arg.addrspace == AddrSpace.GLOBAL: - buf = UOp(Ops.BUFFER, x.dtype, (shape_to_shape_arg((size,)),), ParamArg(next(ctx), device=x.arg.device, addrspace=AddrSpace.GLOBAL)) + buf = UOp(Ops.BUFFER, src=(shape_to_shape_arg((size,)),), arg=ParamArg(next(ctx), x.dtype, device=x.arg.device, addrspace=AddrSpace.GLOBAL)) if x.src[0].op is Ops.SLICE: # no INDEX on SLICE, this could be cleaner do_store = buf.store(x.src[0]).end(*rngs) @@ -439,8 +439,8 @@ class LocalAddBufferContext: opts:tuple|None = None def debuf(ctx:LocalAddBufferContext, buf:UOp): - param = UOp(Ops.PARAM, buf.dtype, (UOp.const(dtypes.int, prod(buf.max_shape)),), - arg=ParamArg(ctx.dg, addrspace=buf.addrspace, device=buf.device)) + param = UOp(Ops.PARAM, src=(UOp.const(dtypes.int, prod(buf.max_shape)),), + arg=ParamArg(ctx.dg, buf.dtype, addrspace=buf.addrspace, device=buf.device)) ret = param.reshape(buf.max_shape) # if the buffer has symbolic shape, shrink the max-sized view to the actual shape if buf.max_shape != buf.shape: ret = ret.shrink(tuple((0, s) for s in buf.shape)) diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index 2eeabecfee..df7a8ba6b8 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -1,7 +1,7 @@ from __future__ import annotations from typing import Any, Callable, cast, TYPE_CHECKING, Type, Sequence, Iterable, Final, Iterator import sys, time, functools, itertools, math, operator, hashlib, os, types, pickle, pathlib, inspect, weakref, collections, struct -from dataclasses import dataclass +from dataclasses import dataclass, replace from enum import Enum, auto from tinygrad.uop import Ops, GroupOp from tinygrad.dtype import ConstType, dtypes, DType, DTypeLike, to_dtype, truncate, least_upper_dtype, Invalid, AddrSpace @@ -22,6 +22,7 @@ class AxisType(Enum): @dataclass(frozen=True, order=True) class ParamArg: slot: int + dtype: DType vmin_vmax: tuple[PyConst, PyConst]|None = None name: str|None = None addrspace: AddrSpace|None = AddrSpace.GLOBAL @@ -29,7 +30,7 @@ class ParamArg: device: str|tuple[str, ...]|None = None def __repr__(self): fields = (("vmin_vmax", None), ("name", None), ("addrspace", AddrSpace.GLOBAL), ("axis", None), ("device", None)) - args = [repr(self.slot)] + [f"{k}={v!r}" for k,default in fields if (v:=getattr(self, k)) != default] + args = [repr(self.slot), repr(self.dtype)] + [f"{k}={v!r}" for k,default in fields if (v:=getattr(self, k)) != default] return f"ParamArg({', '.join(args)})" axis_letters = {AxisType.GLOBAL: "g", AxisType.THREAD: "t", AxisType.LOCAL: "l", AxisType.WARP: "w", AxisType.LOOP: "L", AxisType.UPCAST: "u", AxisType.GROUP_REDUCE: "G", AxisType.REDUCE: "R", AxisType.UNROLL: "r"} @@ -139,15 +140,14 @@ def dtype_from_uop(op:Ops, src:tuple[UOp,...], arg:Any) -> DType|None: assert dtypes.is_int(src[1].dtype), "shift distance must be int" return src[0].dtype case Ops.BUFFER | Ops.PARAM: - # TODO: dtype should move to ParamArg assert isinstance(arg, ParamArg), "BUFFER/PARAM must have ParamArg" - return None + return arg.dtype case Ops.SLICE: # TODO: slice just shouldn't exist return None case Ops.CAST | Ops.BITCAST: - # TODO: dtype should move to arg - return None + assert isinstance(arg, DType), f"CAST/BITCAST arg must be DType, got {arg}" + return arg case Ops.CONST: # TODO: need const refactor to bool/weakint/weakfloat return None @@ -561,10 +561,10 @@ class UOp(RandMixin, metaclass=UOpMetaClass): def cast(self, dtype:DTypeLike): dtype = to_dtype(dtype) if self.dtype == dtype: return self - return UOp(Ops.CAST, dtype, (self,)) + return UOp(Ops.CAST, arg=dtype, src=(self,)) def bitcast(self, dtype:DTypeLike): dtype = to_dtype(dtype) - return self if self.dtype == dtype else UOp(Ops.BITCAST, dtype, (self,)) + return self if self.dtype == dtype else UOp(Ops.BITCAST, arg=dtype, src=(self,)) def load(self, *src:UOp, **kwargs): return UOp(Ops.LOAD, src=(self,)+src, **kwargs) def store(self, src:UOp|ConstType, gate:UOp|None=None, **kwargs): srcs = (self, self.const_like(src) if not isinstance(src, UOp) else src) + ((gate,) if gate is not None else ()) @@ -763,7 +763,7 @@ class UOp(RandMixin, metaclass=UOpMetaClass): @staticmethod def new_buffer(device:str|tuple[str, ...], size:int, dtype:DType, num=None): slot = next(UOp.unique_num) if num is None else num - return UOp(Ops.BUFFER, dtype, (shape_to_shape_arg((size,)),), ParamArg(slot, device=device)) + return UOp(Ops.BUFFER, src=(shape_to_shape_arg((size,)),), arg=ParamArg(slot, dtype, device=device)) @staticmethod def from_buffer(opaque:Buffer, device:str|tuple[str, ...]|None=None): if (uop:=UOp.new_buffer(device or opaque.device, opaque.size, opaque.dtype, num=-id(opaque))) not in buffers: buffers[uop] = opaque.ref(1) @@ -919,8 +919,8 @@ class UOp(RandMixin, metaclass=UOpMetaClass): @staticmethod def variable(name:str, min_val:PyConst, max_val:PyConst, dtype:DType=dtypes.index) -> UOp: - return UOp(Ops.PARAM, dtype, src=(shape_to_shape_arg(()),), - arg=ParamArg(-1, name=name, vmin_vmax=(min_val, max_val), addrspace=AddrSpace.ALU)) + return UOp(Ops.PARAM, src=(shape_to_shape_arg(()),), + arg=ParamArg(-1, dtype, name=name, vmin_vmax=(min_val, max_val), addrspace=AddrSpace.ALU)) @property def expr(self) -> str: assert self.op is Ops.PARAM @@ -1070,11 +1070,11 @@ class UOp(RandMixin, metaclass=UOpMetaClass): @staticmethod def placeholder(shape:tuple[int, ...], dtype:DType, slot:int, addrspace=AddrSpace.GLOBAL): if addrspace is AddrSpace.GLOBAL: - ret = UOp(Ops.PARAM, dtype, src=(shape_to_shape_arg((prod(shape),)),), arg=ParamArg(slot, addrspace=addrspace)) + ret = UOp(Ops.PARAM, src=(shape_to_shape_arg((prod(shape),)),), arg=ParamArg(slot, dtype, addrspace=addrspace)) else: assert addrspace in (AddrSpace.LOCAL, AddrSpace.REG) buf_shape = (prod(shape),) - ret = UOp(Ops.BUFFER, dtype, src=(shape_to_shape_arg(buf_shape),), arg=ParamArg(slot, addrspace=addrspace)) + ret = UOp(Ops.BUFFER, src=(shape_to_shape_arg(buf_shape),), arg=ParamArg(slot, dtype, addrspace=addrspace)) if len(shape) > 1: ret = ret.reshape(shape) return ret def placeholder_like(self, slot:int, addrspace=AddrSpace.GLOBAL): @@ -1092,7 +1092,7 @@ class UOp(RandMixin, metaclass=UOpMetaClass): if shape is not None and axis is not None and isinstance(device, tuple): shape = tuple(s*len(device) if i == axis else s for i,s in enumerate(shape)) src: tuple[UOp, ...] = (UOp(Ops.NOOP) if shape is None else shape_to_shape_arg(shape),) - return UOp(Ops.PARAM, dtype, src, arg=ParamArg(slot, vmin_vmax, name, addrspace, axis, device)) + return UOp(Ops.PARAM, src=src, arg=ParamArg(slot, dtype, vmin_vmax, name, addrspace, axis, device)) def param_like(self, slot:int): addrspace = self.addrspace if self.addrspace is not None else AddrSpace.GLOBAL if self.op is Ops.BIND: @@ -1684,7 +1684,7 @@ pm_lower_index_dtype = PatternMatcher([ (UPat(Ops.SPECIAL, src=(UPat.var("var").cast(dtypes.index),), name="u"), lambda u,var: u.replace(dtype=dtypes.int, src=(var,)).cast(dtypes.index)), (UPat(Ops.PARAM, dtype=dtypes.index, name="u"), - lambda u: u.replace(dtype=dtypes.int).cast(dtypes.index) if u.addrspace == AddrSpace.ALU else None), + lambda u: u.replace(dtype=None, arg=replace(u.arg, dtype=dtypes.int)).cast(dtypes.index) if u.addrspace == AddrSpace.ALU else None), (UPat(Ops.BIND, src=(UPat.var("var").cast(dtypes.index), UPat.cvar("val").cast(dtypes.index))), lambda var,val: var.bind(val).cast(dtypes.index)), # remove hanging casts diff --git a/tinygrad/uop/spec.py b/tinygrad/uop/spec.py index e46535a3ff..a618e1740f 100644 --- a/tinygrad/uop/spec.py +++ b/tinygrad/uop/spec.py @@ -72,7 +72,7 @@ spec_shared = PatternMatcher([ (UPat(GroupOp.ALU, name="x"), lambda x: all(x.dtype == y.dtype for y in x.src)), # CAST - (UPat((Ops.BITCAST, Ops.CAST), src=(UPat(),), name="x"), lambda x: x.arg is None), + (UPat((Ops.BITCAST, Ops.CAST), src=(UPat(),), name="x"), lambda x: isinstance(x.arg, DType)), # RANGE can be in the big graph now (UPat(Ops.RANGE, src=(UPat.var("x"),), allow_any_len=True, name="rng"), lambda rng,x: