diff --git a/tinygrad/codegen/opt/postrange.py b/tinygrad/codegen/opt/postrange.py index aa2e562043..cda176f7e7 100644 --- a/tinygrad/codegen/opt/postrange.py +++ b/tinygrad/codegen/opt/postrange.py @@ -193,7 +193,7 @@ class Scheduler: for b in self.bufs: if rng in (i:=b.src[1].get_idx()).backward_slice_with_self: nb = b.replace(src=(b.src[0], i.valid(valid&b.src[1].get_valid()))) - replaces[b] = nb if b in store_targets else valid.where(nb, UOp.const(Invalid, b.dtype)) + replaces[b] = nb if b in store_targets else valid.where(nb, UOp.const(Invalid)) self.ast = self.ast.substitute(replaces, f"padto {rng.arg[:-1]} {opt.arg}") elif opt.op is OptOps.SWAP: try: diff --git a/tinygrad/llm/kernels/amd.py b/tinygrad/llm/kernels/amd.py index 65c0b68158..b206c3415f 100644 --- a/tinygrad/llm/kernels/amd.py +++ b/tinygrad/llm/kernels/amd.py @@ -86,8 +86,8 @@ def _amd_load(ptr:UOp, lanes:int|None=None) -> UOp: assert ptr.op is Ops.INDEX if lanes is None: return UOp(Ops.CUSTOMI, ptr.dtype, (ptr,), arg="__builtin_nontemporal_load({0})") buf, coords = ptr.src[0], ptr.src[1:] - idx = sum((coord*math.prod(buf.shape[i+1:]) for i,coord in enumerate(coords)), UOp.const(0, dtypes.weakint)) - return UOp(Ops.SHRINK, src=(buf.flatten(), idx, UOp.const(lanes, dtypes.weakint))).load(dtype=ptr.dtype) + idx = sum((coord*math.prod(buf.shape[i+1:]) for i,coord in enumerate(coords)), UOp.const(0)) + return UOp(Ops.SHRINK, src=(buf.flatten(), idx, UOp.const(lanes))).load(dtype=ptr.dtype) def _load_byte(raw:UOp, base:UOp, offset:UOp) -> UOp: return (raw[base + offset//4] >> ((offset&3)*8).cast(dtypes.uint32)) & 255 def _half(value:UOp) -> UOp: return value.cast(dtypes.uint16).bitcast(dtypes.float16).float() diff --git a/tinygrad/runtime/ops_dsp.py b/tinygrad/runtime/ops_dsp.py index 0f001ed79c..e487f666e7 100644 --- a/tinygrad/runtime/ops_dsp.py +++ b/tinygrad/runtime/ops_dsp.py @@ -9,18 +9,10 @@ from tinygrad.renderer.cstyle import ClangRenderer from tinygrad.runtime.autogen import libc, qcom_dsp if getenv("IOCTL"): import extra.dsp.run # noqa: F401 # pylint: disable=unused-import -from tinygrad.uop.ops import PatternMatcher, UPat - -# NOTE: this just increases readability of the generated code -dsp_string = PatternMatcher([ - (UPat(Ops.CONST, (dtypes.int8, dtypes.uint8), name="x"), lambda ctx,x: str(x.val)), -]) - class DSPRenderer(ClangRenderer): has_threads = False buffer_suffix = " restrict __attribute__((align_value(128)))" kernel_typedef = "__attribute__((noinline)) void" - string_rewrite = dsp_string+ClangRenderer.string_rewrite type_map = { **ClangRenderer.type_map, dtypes.uint64: "unsigned long long", dtypes.int64: "long long" } code_for_op = {k:v for k,v in ClangRenderer.code_for_op.items() if k != Ops.SQRT} diff --git a/tinygrad/uop/render.py b/tinygrad/uop/render.py index 1028f524c4..7f070a06a7 100644 --- a/tinygrad/uop/render.py +++ b/tinygrad/uop/render.py @@ -40,7 +40,7 @@ renderer = PatternMatcher([ (UPat(Ops.RANGE, name="x"), lambda x: f"r{range_str(x)}"), (UPat(Ops.CONST, name="x"), lambda x: str(x.val)), # CAST states the width, the weak CONST carries the value - (UPat.cvar("c", dtypes.weaks+(dtypes.bool,)).cast(), lambda c: str(c.val)), + (UPat.cvar("c").cast(), lambda c: str(c.val)), (UPat(Ops.CAST, name="x"), lambda ctx,x: f"({str(x.dtype)[7:]})({ctx[x.src[0]]})"), (UPat(Ops.NEG, name="x"), lambda ctx,x: f"(-{ctx[x.src[0]]})"), (UPat(Ops.RECIPROCAL, name="x"), lambda ctx,x: f"(1/{ctx[x.src[0]]})"), @@ -83,9 +83,9 @@ def render_marg(ctx,x:UOp): sugar = {Ops.SINK, Ops.END, Ops.STORE, Ops.LOAD, Ops.SQRT, Ops.INDEX, Ops.REDUCE, Ops.AFTER, Ops.THREEFRY, Ops.RECIPROCAL, Ops.EXP2, Ops.LOG2, Ops.SIN, Ops.CONTIGUOUS, Ops.BARRIER, Ops.DETACH} pm_pyrender_extra = PatternMatcher([ - (UPat(Ops.CONST, src=(), name="x"), lambda x: f"UOp.const({x.val}, {x.dtype})"), + (UPat(Ops.CONST, src=(), name="x"), lambda x: f"UOp.const({x.val})"), (UPat((Ops.CAST, Ops.BITCAST), name="x"), lambda ctx,x: f"{ctx[x.src[0]]}.{x.op.name.lower()}({x.dtype})" if x.dtype != x.src[0].dtype else None), - (UPat(Ops.SPECIAL, src=(UPat(Ops.CONST),), name="x"), lambda x: f"UOp.special({x.src[0].val}, {repr(x.arg)}, dtype={x.dtype})"), + (UPat(Ops.SPECIAL, src=(UPat(Ops.CONST),), name="x"), lambda x: f"UOp.special({x.src[0].val}, {repr(x.arg)})"), (UPat(Ops.BUFFER, src=(UPat(),), name="x"), lambda x: f"UOp.new_buffer({repr(x.arg.device)}, {x.max_numel()}, {x.dtype}, {x.arg.slot})" if isinstance(x.arg, ParamArg) and x.addrspace is AddrSpace.GLOBAL else None), @@ -95,7 +95,7 @@ pm_pyrender_extra = PatternMatcher([ # NOTE: range has srcs sometimes after control flow (UPat(Ops.RANGE, src=(UPat(Ops.CONST, name="c"),), allow_any_len=True, name="x"), lambda ctx,x,c: "UOp.range("+', '.join([str(c.val)] + [repr(y) for y in x.arg])+ - (f', src={srcs(ctx, x.src[1:])}' if len(x.src) > 1 else '')+(', dtype='+str(x.dtype) if x.dtype is not dtypes.weakint else '')+")"), + (f', src={srcs(ctx, x.src[1:])}' if len(x.src) > 1 else '')+")"), # TODO: index shouldn't mismatch dtype (UPat(Ops.INDEX, src=(UPat(), UPat()), allow_any_len=True, name="x"), lambda ctx,x: f"{ctx[x.src[0]]}.index({ctx[x.src[1]]}, "+''.join([f"{ctx[xx]}, " for xx in x.src[2:]])+ diff --git a/tinygrad/uop/validate.py b/tinygrad/uop/validate.py index 05c8cd3b01..a408c28e4c 100644 --- a/tinygrad/uop/validate.py +++ b/tinygrad/uop/validate.py @@ -45,7 +45,7 @@ z3_renderer = PatternMatcher([ (UPat((Ops.LOAD, Ops.INDEX), dtypes.bool), lambda ctx: (z3.Bool(f"load{len(ctx[1])}", ctx=ctx[0]), None)), # constants (UPat(Ops.CONST, arg=Invalid), lambda ctx: (z3.Int("Invalid", ctx=ctx[0]), None)), - (UPat(Ops.CONST, dtypes.ints+(dtypes.weakint,), name="x"), lambda x,ctx: (z3.IntVal(x.val, ctx=ctx[0]), None)), + (UPat(Ops.CONST, dtypes.weakint, name="x"), lambda x,ctx: (z3.IntVal(x.val, ctx=ctx[0]), None)), (UPat(Ops.CONST, dtypes.bool, name="x"), lambda x,ctx: (z3.BoolVal(x.val, ctx=ctx[0]), None)), # casts from floats create new variables (UPat(Ops.CAST, dtypes.ints+(dtypes.weakint,), src=(UPat(dtype=dtypes.floats),), name="x"), lambda x,ctx: