diff --git a/test/unit/test_graph_rewrite.py b/test/unit/test_graph_rewrite.py index 8a3496101c..30d7454210 100644 --- a/test/unit/test_graph_rewrite.py +++ b/test/unit/test_graph_rewrite.py @@ -303,8 +303,8 @@ class TestRecurse(unittest.TestCase): def test_inf_loop(self): a = UOp.variable('a', 0, 10) pm = PatternMatcher([ - (UPat(Ops.DEFINE_VAR, name="x"), lambda x: x.replace(op=Ops.DEFINE_REG)), - (UPat(Ops.DEFINE_REG, name="x"), lambda x: x.replace(op=Ops.DEFINE_VAR)), + (UPat(Ops.DEFINE_VAR, name="x"), lambda x: x.replace(op=Ops.CONST)), + (UPat(Ops.CONST, name="x"), lambda x: x.replace(op=Ops.DEFINE_VAR)), ]) with self.assertRaises(RuntimeError): graph_rewrite(a, pm) @@ -312,8 +312,8 @@ class TestRecurse(unittest.TestCase): def test_inf_loop_bottom_up(self): a = UOp.variable('a', 0, 10) pm = PatternMatcher([ - (UPat(Ops.DEFINE_VAR, name="x"), lambda x: x.replace(op=Ops.DEFINE_REG)), - (UPat(Ops.DEFINE_REG, name="x"), lambda x: x.replace(op=Ops.DEFINE_VAR)), + (UPat(Ops.DEFINE_VAR, name="x"), lambda x: x.replace(op=Ops.CONST)), + (UPat(Ops.CONST, name="x"), lambda x: x.replace(op=Ops.DEFINE_VAR)), ]) with self.assertRaises(RuntimeError): graph_rewrite(a, pm, bottom_up=True) diff --git a/test/unit/test_viz.py b/test/unit/test_viz.py index 31aa4e1134..020a8e465b 100644 --- a/test/unit/test_viz.py +++ b/test/unit/test_viz.py @@ -124,10 +124,10 @@ class TestViz(BaseTestViz): def test_inf_loop(self): a = UOp.variable('a', 0, 10) - b = a.replace(op=Ops.DEFINE_REG) + b = a.replace(op=Ops.CONST) pm = PatternMatcher([ - (UPat(Ops.DEFINE_VAR, name="x"), lambda x: x.replace(op=Ops.DEFINE_REG)), - (UPat(Ops.DEFINE_REG, name="x"), lambda x: x.replace(op=Ops.DEFINE_VAR)), + (UPat(Ops.DEFINE_VAR, name="x"), lambda x: x.replace(op=Ops.CONST)), + (UPat(Ops.CONST, name="x"), lambda x: x.replace(op=Ops.DEFINE_VAR)), ]) with self.assertRaises(RuntimeError): exec_rewrite(a, [pm]) graphs = flatten(x["graph"].values() for x in get_details(tracked_ctxs[0][0]))