diff --git a/tinygrad/renderer/cstyle.py b/tinygrad/renderer/cstyle.py index 2140abe6e7..e6d01bfc97 100644 --- a/tinygrad/renderer/cstyle.py +++ b/tinygrad/renderer/cstyle.py @@ -1,7 +1,7 @@ from typing import Literal, Callable, cast import os, math, sys from collections import defaultdict, Counter -from tinygrad.codegen.opt import tc +from tinygrad.codegen.opt import tc, axis_letters from tinygrad.uop.ops import GroupOp, Ops, UOp, PatternMatcher, UPat, range_str from tinygrad.helpers import strip_parens, getenv, prod, dedup, AMX, CPU_COUNT from tinygrad.dtype import ImageDType, dtypes, DType, PtrDType, AddrSpace, truncate @@ -163,7 +163,7 @@ class CStyleLanguage(Renderer): # naming prefix = None if u.op is Ops.SPECIAL: r[u] = u.arg - elif u.op is Ops.RANGE: r[u] = "ridx"+range_str(u) + elif u.op is Ops.RANGE: r[u] = f"{axis_letters[u.arg[-1]]}idx"+range_str(u) else: prefix = {Ops.WMMA: "wmma", Ops.DEFINE_LOCAL: "temp", Ops.CONST: "const", Ops.CAST: "cast", Ops.BITCAST: "cast", Ops.GEP: "gep", Ops.VECTORIZE: "cast", Ops.PRECAST: "precast", diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index 951bc4ce8d..322d713852 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -169,7 +169,7 @@ class UOp(MathTrait, metaclass=UOpMetaClass): @property def ptrdtype(self) -> PtrDType: - if not isinstance(self.dtype, PtrDType): raise RuntimeError("ptrdtype called on UOp without PtrDType") + if not isinstance(self.dtype, PtrDType): raise RuntimeError(f"ptrdtype called on UOp with type {self.dtype}") return self.dtype # *** uop shape stuff *** @@ -1178,8 +1178,6 @@ pm_lower_index_dtype = PatternMatcher([ lambda s: s.replace(src=s.src[:2]+tuple(u.src[0] for u in s.src[2:]))), # TODO: this is only triggering if they are all casts, correct? (UPat((Ops.SINK, Ops.NOOP), src=UPat().cast(dtypes.index), name="n"), lambda n: n.replace(src=tuple(s.src[0] for s in n.src))), - # TODO: this should be more general - (UPat(Ops.AFTER, name="x"), lambda x: x.replace(src=tuple(y.src[0] if y.op is Ops.CAST and y.dtype.scalar()==dtypes.index else y for y in x.src))), ]) def _index_to_concrete_int(u:UOp): return graph_rewrite(u.sink(), pm_lower_index_dtype).src[0] diff --git a/tinygrad/uop/symbolic.py b/tinygrad/uop/symbolic.py index e9fec2ae9e..91cf8390e4 100644 --- a/tinygrad/uop/symbolic.py +++ b/tinygrad/uop/symbolic.py @@ -377,6 +377,11 @@ symbolic = symbolic_simple+commutative+PatternMatcher([ (UPat(GroupOp.Binary, src=(UPat.var("x", dtypes.long), UPat.var("y", dtypes.long)), name="u"), lambda u,x,y: x.cast(dtypes.int).alu(u.op, y.cast(dtypes.int)).cast(u.dtype) if not any(v.overflows(dtypes.int) for v in (u,x,y)) else None), ((UPat.var("x", dtypes.index) + UPat.cvar("c")).cast(dtypes.sints, name="cast"), lambda x,c,cast:x.cast(cast.dtype)+c.cast(cast.dtype)), + # only RANGE/IF/STORE/KERNEL have side effects + (UPat(Ops.AFTER, name="x"), lambda x: x.replace(src=(x.src[0],)+ + tuple(flatten([(y,) if y.op in {Ops.RANGE, Ops.IF, Ops.STORE, Ops.KERNEL, Ops.BARRIER} else y.src for y in x.src[1:]])))), + # after with 1 src is just src[0] + (UPat(Ops.AFTER, src=(UPat.var("s"),)), lambda s: s), ])+gep_pushing symbolic_flat = symbolic+PatternMatcher([