fix: minor tweaks to view (#1842)

Co-authored-by: Roelof van Dijk <[email protected]>
This commit is contained in:
Roelof van Dijk
2023-09-10 15:55:57 -07:00
committed by GitHub
co-authored by Roelof van Dijk
parent 47e602f717
commit 1bc52c60df
2 changed files with 3 additions and 8 deletions
+1 -1
View File
@@ -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]
+2 -7
View File
@@ -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)