diff --git a/test/backend/test_const_folding.py b/test/backend/test_const_folding.py index 0c2c801560..fc8ad9f0a8 100644 --- a/test/backend/test_const_folding.py +++ b/test/backend/test_const_folding.py @@ -1,7 +1,7 @@ import unittest, math from tinygrad import Tensor, Device, dtypes from tinygrad.dtype import DTYPES_DICT -from tinygrad.uop.ops import Ops, UOp +from tinygrad.uop.ops import Ops, UOp, GroupOp from tinygrad.codegen.decomp.op import threefry2x32 import numpy as np from test.helpers import not_support_multi_device @@ -17,7 +17,7 @@ def _check_ast_count(desired_count:int, t:Tensor): class TestMovedConstFolding(unittest.TestCase): def test_contiguous_deviceless_const(self): t = Tensor(UOp.const(2.0, dtypes.float)).contiguous() - self.assertIs(t.uop.op, Ops.CONST) + self.assertIs(t.uop, UOp.const(2.0, dtypes.float)) self.assertIsNone(t.uop.device) def test_add_shrunk_zero(self): @@ -169,8 +169,8 @@ class TestMultiConstFolding(unittest.TestCase): class TestThreefryConstFolding(unittest.TestCase): def test_threefry(self): # THREEFRY(const,const) folds to a const once decomposed - x = threefry2x32(UOp.const(5, dtypes.uint64), UOp.const(10, dtypes.uint64)) - self.assertIs(x.simplify().op, Ops.CONST) + x = threefry2x32(UOp.const(5, dtypes.uint64), UOp.const(10, dtypes.uint64)).simplify() + self.assertEqual([u.op for u in x.toposort() if u.op in GroupOp.ALU], []) class TestTautologicalCompare(unittest.TestCase): # without const folding, these would have triggered -Wtautological-compare in clang diff --git a/test/backend/test_linearizer.py b/test/backend/test_linearizer.py index b21dc7ace6..2d527140cb 100644 --- a/test/backend/test_linearizer.py +++ b/test/backend/test_linearizer.py @@ -16,8 +16,6 @@ from test.helpers import replace_opts, check_schedule from test.backend.test_softmax_fusion import single_kernel_softmax MOCKGPU = DEV.interface.startswith("MOCK") -from tinygrad.uop.render import print_uops # noqa: F401 # pylint: disable=unused-import - @unittest.skipIf(isinstance(Device[Device.DEFAULT].renderer, ISARenderer), "isa backends don't preserve the op spec when lowering") class TestLinearizer(unittest.TestCase): def test_arg_dedup(self): @@ -248,7 +246,6 @@ class TestLinearizer(unittest.TestCase): uops = tuple(to_program(replace_opts(ast, opt), renderer=Device[Device.DEFAULT].renderer).src[1].src) begin_range = [i for i, x in enumerate(uops) if x.op is Ops.RANGE][-1] end_range = [i for i, x in enumerate(uops) if x.op is Ops.END][0] - for i,u in enumerate(uops): print(i, u.op, [uops.index(s) for s in u.src], u.arg, u.dtype) for u in uops: if u.op is Ops.STORE and u.src[0].addrspace is AddrSpace.REG: if uops.index(u) < begin_range: diff --git a/test/null/test_uop_graph.py b/test/null/test_uop_graph.py index 180ee23189..e144da641a 100644 --- a/test/null/test_uop_graph.py +++ b/test/null/test_uop_graph.py @@ -13,25 +13,17 @@ simple_pm = PatternMatcher([ ((UPat.var('x') + UPat.cvar('c1')) + UPat.cvar('c2'), lambda x,c1,c2: x + (c1.val+c2.val)), ]) -def const_values(u:UOp): - if u.op is Ops.CONST: return (u.val,) - if u.op is Ops.STACK: return tuple(x.val for x in u.src) - raise AssertionError(f"expected const-like UOp, got {u.op}") - class TestGraphRewriteConst(unittest.TestCase): def test_gep_const(self): v1 = UOp.const((0,1,2), dtypes.int) v2 = v1.index(1) ret = graph_rewrite(v2, sym) - self.assertEqual(ret.dtype, dtypes.int) - self.assertEqual(ret.val, 1) + self.assertIs(ret, UOp.const(1, dtypes.int)) def test_add_const(self): v1 = UOp.const((0,1,2)) v2 = UOp.const((5,6,7)) - ret = graph_rewrite(v1+v2, sym) - self.assertEqual(ret.op, Ops.STACK) - self.assertEqual(const_values(ret), (5,7,9)) + self.assertIs(graph_rewrite(v1+v2, sym), UOp.const((5,7,9))) def xfail_broken_const_wraparound(fn): fn = pytest.mark.xfail(reason="const folding does not properly implement modular arithmetic")(fn) diff --git a/test/null/test_uops.py b/test/null/test_uops.py index 40115c3e44..de18f1a931 100644 --- a/test/null/test_uops.py +++ b/test/null/test_uops.py @@ -129,7 +129,7 @@ class TestConstFloatEq(unittest.TestCase): self.assertFalse(Invalid != HoldsInvalid()) def test_matchers_agree_on_nan(self): - n = UOp.const(math.nan, dtypes.float32) + n = UOp.const(math.nan) for compiled in (False, True): pm = PatternMatcher([(UPat(Ops.CONST, arg=math.nan), lambda: True)], compiled=compiled) self.assertTrue(pm.rewrite(n), f"{compiled=}")