diff --git a/test/null/test_uops.py b/test/null/test_uops.py index de18f1a931..a21fae30a8 100644 --- a/test/null/test_uops.py +++ b/test/null/test_uops.py @@ -5,7 +5,7 @@ from tinygrad.tensor import Tensor from tinygrad.helpers import Timing, Context, cdiv from tinygrad.dtype import dtypes, AddrSpace, ConstFloat, Invalid # noqa: F401 from tinygrad.device import Device -from tinygrad.uop.ops import Ops, ParamArg, PatternMatcher, UOp, UPat, dtype_from_uop, exec_alu, graph_rewrite # noqa: F401 # ParamArg used by eval(str(uop)) roundtrip tests +from tinygrad.uop.ops import Ops, AxisType, ParamArg, PatternMatcher, UOp, UPat, dtype_from_uop, exec_alu, graph_rewrite # noqa: F401 # ParamArg used by eval(str(uop)) roundtrip tests from tinygrad.uop.weak import pm_lower_index_dtype from tinygrad.uop.spec import spec_program, spec_shared, type_verify from tinygrad.uop.symbolic import sym, pm_remove_invalid @@ -457,6 +457,14 @@ class TestUopsObject(unittest.TestCase): self.assertEqual(a.device, Device.DEFAULT) class TestUOpRender(unittest.TestCase): + def test_render_ssimplified_marg_outside_toposort(self): + r = UOp.range(UOp.const(16, dtypes.int), 2, AxisType.WEAK, dtype=dtypes.int) + offset = (r * 2) + (r * 2) + shrink = UOp(Ops.SHRINK, src=(UOp.param(0, dtypes.uint, (32,)), offset, UOp.const(2, dtypes.int))) + self.assertIsNot(shrink.src[1], shrink.marg[0][0]) + self.assertEqual(shrink.render(simplify=False), "p0.shrink((((r2*4), 2),))") + self.assertEqual(UOp.range(1, 0, src=(shrink,), dtype=dtypes.int).render(simplify=False), "r0") + def test_render_vectorize_empty(self): u = UOp(Ops.STACK, dtype=dtypes.void, src=()) self.assertEqual(u.render(simplify=False), "{}") diff --git a/tinygrad/uop/render.py b/tinygrad/uop/render.py index b005a6eaf2..1028f524c4 100644 --- a/tinygrad/uop/render.py +++ b/tinygrad/uop/render.py @@ -1,6 +1,6 @@ from tinygrad.dtype import AddrSpace, dtypes from tinygrad.uop import Ops, GroupOp -from tinygrad.uop.ops import ParamArg, UOp, PatternMatcher, UPat, multirange_str, range_str, consumer_map_from_toposort +from tinygrad.uop.ops import ParamArg, UOp, PatternMatcher, UPat, multirange_str, range_str, consumer_map_from_toposort, sint from tinygrad.helpers import strip_parens def pretty_print(x:UOp, cache=None, d=0)->str: @@ -69,14 +69,15 @@ renderer_infer = PatternMatcher([ # *** pyrender *** def srcs(ctx, src): return f"({ctx[src[0]]},)" if len(src) == 1 else f"({', '.join([ctx[x] for x in src])})" +# marg is ssimplify'd, so a bound can be a node this graph never contained +def marg_str(ctx, a:sint) -> str: return str(a) if not isinstance(a, UOp) else ctx[a] if a in ctx else a.render() + def render_marg(ctx,x:UOp): if x.op is Ops.PERMUTE: return str(x.marg) if x.op is Ops.FLIP: return str(tuple([i for i,x in enumerate(x.marg) if x])) pieces = [] - if x.op in {Ops.RESHAPE, Ops.EXPAND}: - pieces = [f"{ctx[a] if isinstance(a, UOp) else str(a)}" for a in x.marg] - if x.op in {Ops.PAD, Ops.SHRINK}: - pieces = [f"({ctx[a[0]] if isinstance(a[0], UOp) else str(a[0])}, {ctx[a[1]] if isinstance(a[1], UOp) else str(a[1])})" for a in x.marg] + if x.op in {Ops.RESHAPE, Ops.EXPAND}: pieces = [marg_str(ctx, a) for a in x.marg] + if x.op in {Ops.PAD, Ops.SHRINK}: pieces = [f"({marg_str(ctx, a[0])}, {marg_str(ctx, a[1])})" for a in x.marg] return f"({','.join(pieces)})" if len(pieces) != 1 else f"({pieces[0]},)" sugar = {Ops.SINK, Ops.END, Ops.STORE, Ops.LOAD, Ops.SQRT, Ops.INDEX, Ops.REDUCE, Ops.AFTER, Ops.THREEFRY,