From 678f83e41b268110b8f5cccab93bd19c3213667e Mon Sep 17 00:00:00 2001 From: chenyu Date: Thu, 9 Oct 2025 17:06:10 +0800 Subject: [PATCH] delete ShapeTracker to_valid_uop and substitute [pr] (#12563) --- test/unit/test_shapetracker.py | 4 ++-- tinygrad/shape/shapetracker.py | 4 ---- 2 files changed, 2 insertions(+), 6 deletions(-) diff --git a/test/unit/test_shapetracker.py b/test/unit/test_shapetracker.py index ec7b56a20f..2412ea475f 100644 --- a/test/unit/test_shapetracker.py +++ b/test/unit/test_shapetracker.py @@ -3,14 +3,14 @@ import unittest import numpy as np from tinygrad.dtype import dtypes, Invalid from tinygrad.helpers import prod -from tinygrad.shape.shapetracker import ShapeTracker, View +from tinygrad.shape.shapetracker import ShapeTracker, View, views_to_valid_uop from tinygrad import Variable from tinygrad.uop.ops import UOp, Ops, graph_rewrite from tinygrad.codegen.late.devectorizer import sym from itertools import product def shapetracker_getitem(st:ShapeTracker, val:int): - valid_idx = st.reshape((st.size,)).to_valid_uop([UOp.const(dtypes.int, val)]) + valid_idx = views_to_valid_uop(st.reshape((st.size,)).views, (UOp.const(dtypes.int, val),)) idx, valid = valid_idx.get_idx(), valid_idx.get_valid() idx, valid = graph_rewrite(idx, sym), graph_rewrite(valid, sym) assert idx.op is Ops.CONST and valid.op is Ops.CONST diff --git a/tinygrad/shape/shapetracker.py b/tinygrad/shape/shapetracker.py index 9435b909f9..57c84f9e79 100644 --- a/tinygrad/shape/shapetracker.py +++ b/tinygrad/shape/shapetracker.py @@ -53,9 +53,6 @@ class ShapeTracker: @property def size(self) -> int: return self.views[-1].size() - def to_valid_uop(self, _idxs:list[UOp]|tuple[UOp, ...]|None=None) -> UOp: - return views_to_valid_uop(self.views, tuple(_idxs) if _idxs is not None else None) - def vars(self) -> set[Variable]: return set().union(*[v.vars() for v in self.views]) @property @@ -65,7 +62,6 @@ class ShapeTracker: unbound_views, var_vals = zip(*[v.unbind() for v in self.views]) if all(len(x) == 0 for x in var_vals): return self, {} return ShapeTracker(tuple(unbound_views)), merge_dicts(var_vals) - def substitute(self, dvars:dict[UOp, UOp]): return ShapeTracker(tuple(x.substitute(dvars) for x in self.views)) def real_strides(self, ignore_valid=False) -> tuple[sint|None, ...]: with Context(TRACK_MATCH_STATS=0): return views_to_real_strides(self.views, ignore_valid)