mirror of
https://github.com/tinygrad/tinygrad.git
synced 2026-08-29 14:56:06 +00:00
fix uop swizzle on BUFFER, new tests (#7875)
* fix uop swizzle on BUFFER, new tests * can have view of view
This commit is contained in:
@@ -1864,6 +1864,24 @@ class TestSwizzle(unittest.TestCase):
|
||||
ret = swizzle_rewrite(sink)
|
||||
self.assertEqual(swizzle_cnt(ret), 0)
|
||||
|
||||
def test_non_contiguous_view_simplify(self):
|
||||
st = ShapeTracker(views=(View(shape=(2048, 2048), strides=(1, 2048), offset=0, mask=None, contiguous=False),))
|
||||
a = UOp(Ops.LOAD, dtypes.char, (UOp.new_buffer(Device.DEFAULT, 4194304, dtypes.char), st.to_uop()))
|
||||
ret = swizzle_rewrite(a.view(st))
|
||||
self.assertEqual(ret.st_arg, st+st)
|
||||
|
||||
def test_contiguous_view_simplify(self):
|
||||
base = ShapeTracker.from_shape((32, 32))
|
||||
a = UOp(Ops.LOAD, dtypes.char, (UOp.new_buffer(Device.DEFAULT, base.size, dtypes.char), base.to_uop()))
|
||||
swizzle = a.reshape((64, 16))
|
||||
self.assertEqual(swizzle_cnt(swizzle), 1)
|
||||
ret = swizzle_rewrite(swizzle)
|
||||
self.assertEqual(ret.st_arg, base.reshape((64, 16))) # late rewrite
|
||||
reswizzle = a.reshape((64, 16)).reshape((32, 32))
|
||||
self.assertEqual(swizzle_cnt(reswizzle), 0) # instant rule
|
||||
ret = swizzle_rewrite(reswizzle)
|
||||
self.assertIs(ret, reswizzle)
|
||||
|
||||
def store_val(si:ScheduleItem): return si.ast.src[0].src[2]
|
||||
class TestView(unittest.TestCase):
|
||||
def test_all_masked_out(self):
|
||||
|
||||
+4
-2
@@ -354,10 +354,12 @@ class UOp(MathTrait, metaclass=UOpMetaClass):
|
||||
# *** uop movement ops ***
|
||||
|
||||
@property
|
||||
def base(self) -> UOp: return self.src[0] if self.op is Ops.VIEW and len(self.src) == 1 else self
|
||||
def base(self) -> UOp: return self.src[0] if self.op is Ops.VIEW and len(self.src) == 1 and self.src[0].op is not Ops.BUFFER else self
|
||||
def view(self, st:ShapeTracker) -> UOp:
|
||||
if self.st is None: return self
|
||||
assert self.op is not Ops.STORE, "VIEW of STORE is invalid, STORE is always base"
|
||||
return self if self.st is None or self.st == st else UOp(Ops.VIEW, self.dtype, (self,), st)
|
||||
if st.contiguous and self.base.st == st: return self.base
|
||||
return UOp(Ops.VIEW, self.dtype, (self,), st)
|
||||
def reshape(self, arg:Tuple[sint, ...]) -> UOp: return self.view(unwrap(self.st).reshape(arg))
|
||||
|
||||
# *** uop Buffer stuff ***
|
||||
|
||||
Reference in New Issue
Block a user