forked from tinygrad/tinygrad
multitensor reshape tests
This commit is contained in:
@@ -307,6 +307,53 @@ class TestMultiTensor(unittest.TestCase):
|
||||
# for i, ast in enumerate(asts):
|
||||
# print(f"{i} {ast}")
|
||||
|
||||
def test_reshape_on_axis(self):
|
||||
devices = (d0, d1, d2)
|
||||
|
||||
t0 = Tensor.rand((26, 15, 7)).shard(devices, axis=1)
|
||||
|
||||
# test split and rejoin to the right
|
||||
t1 = t0.reshape((26, 3, 5, 7))
|
||||
t2 = t0.reshape((26, 3, 35))
|
||||
t3 = t1.reshape((26, 15, 7))
|
||||
t4 = t2.reshape((26, 105,))
|
||||
|
||||
for t in [t0, t1, t2, t3, t4]:
|
||||
assert t.lazydata.axis == 1
|
||||
np.testing.assert_allclose(t.numpy().flatten(), t0.numpy().flatten())
|
||||
|
||||
# test shape-one axis
|
||||
t5 = t4.reshape((26, 1, 105))
|
||||
assert t5.lazydata.axis == 2
|
||||
|
||||
# test split and rejoin to the right and reshape to the left
|
||||
t5 = t0.reshape((2, 13, 3, 5, 7))
|
||||
t6 = t0.reshape((13, 2, 3, 7, 5))
|
||||
t7 = t0.reshape((1, 13, 2, 3, 1, 7, 5))
|
||||
np.testing.assert_allclose(t5.numpy().flatten(), t0.numpy().flatten())
|
||||
assert t5.lazydata.axis == 2
|
||||
np.testing.assert_allclose(t6.numpy().flatten(), t0.numpy().flatten())
|
||||
assert t6.lazydata.axis == 2
|
||||
np.testing.assert_allclose(t7.numpy().flatten(), t0.numpy().flatten())
|
||||
assert t7.lazydata.axis == 3
|
||||
|
||||
# test no left join
|
||||
with self.assertRaises((AssertionError, ValueError)):
|
||||
t0.reshape((26*15,7))
|
||||
|
||||
def test_reshape_on_axis_uneven(self):
|
||||
devices = (d0, d1, d2)
|
||||
t0 = Tensor.rand((4, 8, 15)).shard(devices, axis=1)
|
||||
|
||||
# no split axis if uneven
|
||||
with self.assertRaises((AssertionError, ValueError)):
|
||||
t0.reshape((4,4,2,15))
|
||||
|
||||
# ok to split reshape left and right though
|
||||
t1 = t0.reshape(2, 2, 8, 3, 5)
|
||||
np.testing.assert_allclose(t0.numpy().flatten(), t1.numpy().flatten())
|
||||
assert t1.lazydata.axis == 2
|
||||
|
||||
@unittest.skipIf(CI and Device.DEFAULT in {"GPU", "CUDA", "METAL"}, "no GPU CI")
|
||||
class TestShrinkMultiTensorShardedAxis(unittest.TestCase):
|
||||
# shrink a multitensor on sharded axis
|
||||
|
||||
Reference in New Issue
Block a user