forked from tinygrad/tinygrad
Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9de19ef89f | ||
|
|
403fdfcfd4 | ||
|
|
22674798df |
@@ -267,6 +267,8 @@ jobs:
|
||||
run: python -c "from tinygrad import Device; assert Device.DEFAULT == 'CPU', Device.DEFAULT"
|
||||
- name: Run unit tests
|
||||
run: CPU=1 python -m pytest -n=auto test/unit/ --durations=20
|
||||
- name: Check SPEC=1
|
||||
run: SPEC=1 python3 test/test_tiny.py
|
||||
- name: Run targetted tests on NULL backend
|
||||
run: NULL=1 python3 -m unittest test.test_multitensor.TestMultiTensor.test_data_parallel_resnet_train_step test/device/test_null.py
|
||||
# TODO: too slow
|
||||
|
||||
+3
-2
@@ -280,13 +280,14 @@ class TestAssign(unittest.TestCase):
|
||||
b.realize()
|
||||
ba1 = a.uop.base.realized
|
||||
bb1 = b.uop.base.realized
|
||||
with self.assertRaises((RuntimeError, AssertionError)):
|
||||
with self.assert_permuted_assign():
|
||||
a = a.permute(1,0)
|
||||
a += b
|
||||
a.realize()
|
||||
ba2 = a.uop.base.realized
|
||||
assert ba1 != ba2 and ba1 != bb1
|
||||
np.testing.assert_allclose(a.numpy(), np.arange(N*N).reshape((N,N)) + np.arange(N*N).reshape((N,N)).transpose(1,0))
|
||||
# permute and base are the same buffer
|
||||
assert ba1 == ba2 and ba1 != bb1
|
||||
|
||||
def test_post_permuted_assignment(self):
|
||||
a = Tensor(np.arange(N*N, dtype=np.float32)).reshape(N,N)
|
||||
|
||||
+44
-9
@@ -302,18 +302,53 @@ class TestOuterworld(unittest.TestCase):
|
||||
|
||||
from tinygrad.schedule.rangeify import pm_rangeify, RangeifyContext
|
||||
class TestRangeifyPM(unittest.TestCase):
|
||||
@unittest.expectedFailure
|
||||
def test_reshape_match(self):
|
||||
def proc(a:Tensor):
|
||||
sink = a.uop.sink()
|
||||
def setUp(self): self.base = Tensor.empty(10*10).reshape(10, 10).contiguous()
|
||||
def assert_same(self, a, b):
|
||||
def run_pm_rangeify(t:Tensor):
|
||||
sink = t.uop.sink()
|
||||
pm_realize = PatternMatcher([(UPat(Ops.CONTIGUOUS, name="x"), lambda x: x.replace(op=Ops.REALIZE))])
|
||||
sink = graph_rewrite(sink, pm_realize)
|
||||
return graph_rewrite(sink, pm_rangeify, ctx=RangeifyContext())
|
||||
a = Tensor.empty(10*10).reshape(10, 10).contiguous().pad(((0,0),(0,1))).contiguous()
|
||||
b = Tensor.empty(10*10).reshape(10, 10).contiguous().reshape(100).reshape(10, 10).pad(((0,0),(0,1))).contiguous()
|
||||
sink1 = proc(a)
|
||||
sink2 = proc(b)
|
||||
self.assertIs(sink1, sink2)
|
||||
self.assertIs(run_pm_rangeify(a.contiguous()), run_pm_rangeify(b.contiguous()))
|
||||
|
||||
def test_nothing_match(self):
|
||||
a = self.base.pad(((0,0),(0,1)))
|
||||
b = self.base.pad(((0,0),(0,1)))
|
||||
self.assert_same(a, b)
|
||||
|
||||
def test_reshape_match(self):
|
||||
a = self.base
|
||||
b = self.base.reshape(100).reshape(10, 10)
|
||||
self.assert_same(a, b)
|
||||
|
||||
def test_permute_reshape_match(self):
|
||||
a = self.base
|
||||
b = self.base.permute(1,0).reshape(100).reshape(10, 10).permute(1,0)
|
||||
self.assert_same(a, b)
|
||||
|
||||
def test_padded_permute_match(self):
|
||||
a = self.base.pad(((0,0),(0,1)))
|
||||
b = self.base.permute(1,0).pad(((0,1),(0,0))).permute(1,0)
|
||||
self.assert_same(a, b)
|
||||
|
||||
@unittest.expectedFailure
|
||||
def test_padded_reshape_match(self):
|
||||
a = self.base.pad(((0,0),(0,1)))
|
||||
b = self.base.reshape(100).reshape(10, 10).pad(((0,0),(0,1)))
|
||||
self.assert_same(a, b)
|
||||
|
||||
@unittest.expectedFailure
|
||||
def test_padded_permute_reshape_match(self):
|
||||
a = self.base.pad(((0,0),(0,1)))
|
||||
b = self.base.permute(1,0).reshape(100).reshape(10, 10).pad(((0,1),(0,0))).permute(1,0)
|
||||
self.assert_same(a, b)
|
||||
|
||||
# why is this failing?
|
||||
@unittest.expectedFailure
|
||||
def test_cross_pad_match(self):
|
||||
a = self.base.pad(((0,0),(0,1))).pad(((0,1),(0,0)))
|
||||
b = self.base.pad(((0,1),(0,0))).pad(((0,0),(0,1)))
|
||||
self.assert_same(a, b)
|
||||
|
||||
class TestRangeifyEdgeCase(unittest.TestCase):
|
||||
def test_matmul_relu_cat(self):
|
||||
|
||||
@@ -568,5 +568,13 @@ class TestUOpChildren(unittest.TestCase):
|
||||
del c
|
||||
self.assertEqual(len(a.children), 0)
|
||||
|
||||
class TestUOpRender(unittest.TestCase):
|
||||
def test_render_vectorize_same(self):
|
||||
u = UOp(Ops.VECTORIZE, src=(UOp.const(dtypes.int, 0), UOp.const(dtypes.int, 0), UOp.const(dtypes.int, 0)))
|
||||
self.assertEqual(u.render(), "{0, ...}")
|
||||
def test_render_vectorize_different(self):
|
||||
u = UOp(Ops.VECTORIZE, src=(UOp.const(dtypes.int, 0), UOp.const(dtypes.int, 1), UOp.const(dtypes.int, 2)))
|
||||
self.assertEqual(u.render(), "{0,1,2}")
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main(verbosity=2)
|
||||
|
||||
+2
-2
@@ -1042,7 +1042,7 @@ if TRACK_MATCH_STATS or PROFILE:
|
||||
|
||||
# *** simple graph rewrite engine ***
|
||||
|
||||
SENTINEL = UOp(Ops.SENTINEL)
|
||||
with Context(SPEC=0): SENTINEL = UOp(Ops.SENTINEL)
|
||||
class RewriteNotReady(Exception): pass
|
||||
class BottomUpGate(Exception): pass
|
||||
class RewriteContext:
|
||||
@@ -1197,7 +1197,7 @@ renderer = PatternMatcher([
|
||||
(UPat((Ops.INDEX, Ops.BUFFERIZE), name="x"), lambda x:
|
||||
UOp(Ops.NOOP, arg=''.join([f"[{strip_parens(y.arg)}]" for y in x.src[1:]])) if all(y.op is Ops.NOOP for y in x.src[1:]) else None),
|
||||
(UPat(Ops.VECTORIZE, src=UPat(Ops.NOOP), name="x"),
|
||||
lambda x: UOp(Ops.NOOP, arg=f"[{','.join([y.arg for y in x.src])}]" if not all_same(x.src) else f"{len(x.src)}x[{x.src[0].arg}]")),
|
||||
lambda x: UOp(Ops.NOOP, arg=f"{{{','.join([y.arg for y in x.src])}}}" if not all_same(x.src) else f"{{{x.src[0].arg}, ...}}")),
|
||||
])
|
||||
renderer_infer = PatternMatcher([
|
||||
(UPat(Ops.MOD, src=UPat(Ops.NOOP), name="x"), lambda x: UOp(Ops.NOOP, arg=f"cmod({x.src[0].arg}, {x.src[1].arg})")),
|
||||
|
||||
@@ -258,6 +258,9 @@ full_non_rangeify_spec = PatternMatcher([]) if RANGEIFY else PatternMatcher([
|
||||
])
|
||||
|
||||
full_spec = PatternMatcher([
|
||||
# SENTINEL should never be in the graph
|
||||
(UPat(Ops.SENTINEL), lambda: False),
|
||||
|
||||
# Invalid must have type Index
|
||||
(UPat(Ops.CONST, arg=Invalid, name="x"), lambda x: x.dtype.scalar() == dtypes.index),
|
||||
# where on index in rhs position is fine
|
||||
|
||||
Reference in New Issue
Block a user