From eda2824cd8b71417c847c9cdc8b470e6439465b5 Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Wed, 26 Jun 2024 10:18:15 -0700 Subject: [PATCH] freeze uop [run_process_replay] (#5155) --- tinygrad/codegen/uops.py | 12 ++++-------- 1 file changed, 4 insertions(+), 8 deletions(-) diff --git a/tinygrad/codegen/uops.py b/tinygrad/codegen/uops.py index d11b2da115..d9971f1c8e 100644 --- a/tinygrad/codegen/uops.py +++ b/tinygrad/codegen/uops.py @@ -27,7 +27,7 @@ class UOps(Enum): ENDRANGE = auto(); ENDIF = auto() # noqa: E702 def ufix(dtype: Optional[DType], x): return UOp.const(dtype, x) if not isinstance(x, UOp) else x -@dataclass(eq=False) +@dataclass(frozen=True, eq=False) class UOp: op: UOps dtype: Optional[DType] = None @@ -313,7 +313,7 @@ class UOpGraph: run_cnt = 0 while changed: changed = 0 - @functools.lru_cache + @functools.lru_cache(None) def rewrite(u:UOp) -> UOp: nonlocal changed recurse_cnt = 0 @@ -324,12 +324,8 @@ class UOpGraph: up = rewritten recurse_cnt += 1 changed += recurse_cnt - # NOTE: this changes UOp, so we have to delete caches - up.src = tuple(rewrite(x) for x in up.src) - if 'parents' in up.__dict__: delattr(up, 'parents') - if 'cmp_tuple' in up.__dict__: delattr(up, 'cmp_tuple') - # replace with cached nodes - return self.nodes.setdefault(up.tuple(), up) + up = UOp(up.op, up.dtype, tuple(rewrite(x) for x in up.src), up.arg) + return self.nodes.setdefault(up.tuple(), up) # replace with cached nodes sink = rewrite(sink) run_cnt += 1 assert run_cnt < 100, "exceeded 100 rewrite loops!"