From cc5e4e54b87e3e69eb26c128642a331ad18643e7 Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Sun, 15 Jun 2025 12:00:52 -0700 Subject: [PATCH] move type verify to codegen [pr] (#10816) --- tinygrad/codegen/__init__.py | 7 ++++++- tinygrad/codegen/linearize.py | 4 ---- 2 files changed, 6 insertions(+), 5 deletions(-) diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index b688940b4b..2474c49a01 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -3,6 +3,7 @@ import functools from dataclasses import dataclass from tinygrad.helpers import QUANTIZE, DEVECTORIZE, TRANSCENDENTAL from tinygrad.uop.ops import PatternMatcher, graph_rewrite, UOp +from tinygrad.uop.spec import type_verify from tinygrad.renderer import Renderer # import all pattern matchers here @@ -72,4 +73,8 @@ def _get_rewrites_for_renderer(opts:Renderer, linearizer:bool, _QUANTIZE, _DEVEC def full_rewrite_to_sink(sink:UOp, opts:Renderer|None=None, linearizer:bool=False) -> UOp: return apply_rewrites(sink, get_rewrites_for_renderer(opts if opts is not None else Renderer(), linearizer)) -def full_rewrite(sink:UOp, opts:Renderer|None=None) -> list[UOp]: return list(full_rewrite_to_sink(sink, opts, linearizer=True).arg.lst) + +def full_rewrite(sink:UOp, opts:Renderer|None=None) -> list[UOp]: + lst = list(full_rewrite_to_sink(sink, opts, linearizer=True).arg.lst) + if __debug__: type_verify(lst) + return lst diff --git a/tinygrad/codegen/linearize.py b/tinygrad/codegen/linearize.py index 23dfc791c2..d205299e54 100644 --- a/tinygrad/codegen/linearize.py +++ b/tinygrad/codegen/linearize.py @@ -4,7 +4,6 @@ from collections import defaultdict from dataclasses import dataclass, replace from tinygrad.uop.ops import UOp, Ops, PatternMatcher, UPat, GroupOp from tinygrad.helpers import dedup, partition, all_same, flatten, getenv -from tinygrad.uop.spec import type_verify # NOTE: any toposort should be valid here, unlike last time this isn't required, it's just for speed def block_reorder(lst:list[UOp]) -> list[UOp]: @@ -237,9 +236,6 @@ def finalize(sink:UOp) -> UOp: # place the early things lst = sorted(dedup(sink.src), key=lambda x: x.tuplize) + list(sink.arg.lst) - - if __debug__: type_verify(lst) - return UOp(Ops.BLOCKFINAL, arg=BasicBlock(tuple(lst))) pm_finalize = PatternMatcher([(UPat(Ops.BLOCK, name="sink"), finalize)])