forked from tinygrad/tinygrad
This reverts commit f19d8bb7b4.
This commit is contained in:
@@ -355,15 +355,6 @@ class TestMultiTensor(unittest.TestCase):
|
||||
np.testing.assert_allclose(t0.numpy().flatten(), t1.numpy().flatten())
|
||||
assert t1.lazydata.axis == 2
|
||||
|
||||
def test_mlb_assign_change_axis(self):
|
||||
devices = (d0, d1)
|
||||
|
||||
t_none = Tensor.zeros((16, 16)).shard(devices).contiguous().realize()
|
||||
t_zero = Tensor.ones((16, 16)).shard(devices, axis=0)
|
||||
with self.assertRaises(AssertionError):
|
||||
# don't allow assigns that change axes
|
||||
t_none.assign(t_zero)
|
||||
|
||||
@unittest.skipIf(CI and Device.DEFAULT in {"GPU", "CUDA", "METAL"}, "no GPU CI")
|
||||
class TestShrinkMultiTensorShardedAxis(unittest.TestCase):
|
||||
# shrink a multitensor on sharded axis
|
||||
|
||||
@@ -142,8 +142,6 @@ class Tensor:
|
||||
if x.__class__ is not Tensor: x = Tensor(x, device=self.device, dtype=self.dtype)
|
||||
# NOTE: we allow cross device assign
|
||||
assert self.shape == x.shape, f"assign shape mismatch {self.shape} != {x.shape}"
|
||||
if isinstance(self.lazydata, MultiLazyBuffer):
|
||||
assert self.lazydata.axis == x.lazydata.axis
|
||||
assert not x.requires_grad # self requires_grad is okay?
|
||||
if DEBUG >= 4: print(f"assign {self.lazydata} <- {x.lazydata}")
|
||||
if self.dtype == x.dtype and not getenv("DISALLOW_ASSIGN"):
|
||||
|
||||
Reference in New Issue
Block a user