forked from tinygrad/tinygrad
arange stack regression test (#17250)
* simple failing test * commend out stack * a little smaller * work
This commit is contained in:
@@ -19,6 +19,12 @@ class TestArange(unittest.TestCase):
|
||||
self.assertLess(self._get_flops(Tensor.arange(256).clone(), np.arange(256)), 256*4)
|
||||
self.assertLess(self._get_flops(Tensor.arange(2560).clone(), np.arange(2560)), 2560*4)
|
||||
|
||||
def test_cat_complexity(self):
|
||||
x = Tensor.arange(2**10) + Tensor.empty((), dtype=dtypes.uint32)
|
||||
out = x.cat(x).cat(Tensor.empty(1, dtype=dtypes.uint32))
|
||||
linear = compile_linear(out.schedule_linear())
|
||||
self.assertLessEqual(estimate_uop(linear.src[-1]).ops, out.numel()*20)
|
||||
|
||||
@unittest.skipIf(Device.DEFAULT == "CL", "flaky in CI")
|
||||
def test_arange_cumsum(self):
|
||||
np.testing.assert_equal(Tensor.arange(513).cumsum(0).numpy(), np.arange(513).cumsum())
|
||||
|
||||
@@ -123,6 +123,7 @@ class TestUOpsStats(unittest.TestCase):
|
||||
# NOTE; ops also include indexing ops
|
||||
assert expected_ops <= ops and ops <= expected_ops * 2
|
||||
|
||||
@unittest.expectedFailure
|
||||
def test_cat_equal_pieces(self):
|
||||
# concatenating equal-size pieces lowers to STACK: pure data movement, no arithmetic
|
||||
equal = [Tensor.empty(256, 128) for _ in range(4)]
|
||||
|
||||
@@ -719,7 +719,8 @@ class OpMixin(ElementwiseMixin, ReduceMixin):
|
||||
dim = self._resolve_dim(dim)
|
||||
for arg in args: assert arg.ndim==self.ndim and all(ti==ai for i,(ti,ai) in enumerate(zip(self.shape, arg.shape)) if i!=dim)
|
||||
tensors = [self, *args]
|
||||
if all(t.shape[dim] == self.shape[dim] for t in args): return self.stack(*args, dim=dim).flatten(dim, dim+1)
|
||||
# TODO: enable this
|
||||
#if all(t.shape[dim] == self.shape[dim] for t in args): return self.stack(*args, dim=dim).flatten(dim, dim+1)
|
||||
dim_cumsum = list(itertools.accumulate([t.shape[dim] for t in tensors], initial=0))
|
||||
padded = [t.pad(tuple((dim_cumsum[i], dim_cumsum[-1]-dim_cumsum[i+1]) if j==dim else None for j in range(t.ndim))) for i,t in enumerate(tensors)]
|
||||
return padded[0].usum(*padded[1:])
|
||||
|
||||
Reference in New Issue
Block a user