From 1bc52c60dfc198b900b277af82f7c2babc58a7ff Mon Sep 17 00:00:00 2001 From: Roelof van Dijk <3604013+roelofvandijk@users.noreply.github.com> Date: Mon, 11 Sep 2023 00:55:57 +0200 Subject: [PATCH] fix: minor tweaks to view (#1842) Co-authored-by: Roelof van Dijk --- tinygrad/codegen/optimizer.py | 2 +- tinygrad/shape/shapetracker.py | 9 ++------- 2 files changed, 3 insertions(+), 8 deletions(-) diff --git a/tinygrad/codegen/optimizer.py b/tinygrad/codegen/optimizer.py index 89bf2875f0..e56c698c9c 100644 --- a/tinygrad/codegen/optimizer.py +++ b/tinygrad/codegen/optimizer.py @@ -120,7 +120,7 @@ class OptimizedKernel(Kernel): stride[j] = bst bst *= shp[j] - self.sts.append(ShapeTracker(tuple(shp), [View(tuple(shp), tuple(stride))])) + self.sts.append(ShapeTracker(tuple(shp), [View.create(tuple(shp), tuple(stride))])) self.bufs.append(LocalBuffer(name=f"ldata{i}", size=self.sts[-1].size())) if DEBUG >= 4: print("aliasing buffer", self.sts[i]) self.local_alias[i] = self.bufs[-1] diff --git a/tinygrad/shape/shapetracker.py b/tinygrad/shape/shapetracker.py index 56f25f6562..525fbd56de 100644 --- a/tinygrad/shape/shapetracker.py +++ b/tinygrad/shape/shapetracker.py @@ -25,7 +25,6 @@ def is_contiguous(shape:Tuple[int, ...], strides:Tuple[int, ...]) -> bool: retur def filter_strides(shape:Tuple[int, ...], strides:Tuple[int, ...]) -> Tuple[int, ...]: return tuple(stride if shp != 1 else 0 for stride, shp in zip(strides, shape)) -@functools.lru_cache(maxsize=None) class View(NamedTuple): shape:Tuple[int, ...] strides:Tuple[int, ...] @@ -33,16 +32,12 @@ class View(NamedTuple): mask:Optional[Tuple[Tuple[int, int]]] = None @staticmethod + @functools.lru_cache(maxsize=None) def create(shape, strides=None, offset=0, mask=None): return View(shape, filter_strides(shape, strides) if strides else strides_for_shape(shape), offset, mask) - def __repr__(self): return f"View(shape={self.shape}, strides={self.strides}, offset={self.offset}, mask={self.mask})" - @property def contiguous(self): return self.offset == 0 and is_contiguous(self.shape, self.strides) and self.mask is None - @property - def shape_strides(self): return to_shape_strides(self.shape, self.strides) - def expr_node_mask(self, idx, valid=None) -> Node: expr = [valid] if valid is not None else [] if self.mask is not None: @@ -59,7 +54,7 @@ class View(NamedTuple): if idx is None: idx = Variable('idx', 0, prod(self.shape)-1) ret: List[Node] = [Variable.num(self.offset) if isinstance(self.offset, int) else self.offset] if self.offset else [] acc = 1 - for d,s in reversed(self.shape_strides): + for d,s in reversed(to_shape_strides(self.shape, self.strides)): ret.append(((idx//acc)%d)*s) acc *= d return Variable.sum(ret)