diff --git a/test/null/test_uop_graph.py b/test/null/test_uop_graph.py index 4f60284230..24d17b6c2e 100644 --- a/test/null/test_uop_graph.py +++ b/test/null/test_uop_graph.py @@ -248,6 +248,14 @@ class TestUOpGraph(unittest.TestCase): uops = to_uops_list([out]) self.assertEqual(len(uops), 2) # +1 for SINK + def test_devectorize_derives_lane_dtype(self): + from tinygrad.codegen import do_devectorize + # an Invalid lane derives bool while the value lane derives float: the lane rebuild must derive, not inherit + lhs = UOp.stack(UOp.invalid(), UOp.const(None, 1.0).cast(dtypes.float)) + out = do_devectorize(lhs * lhs) + invalid_lane_mul = next(u for u in out.src[0].toposort() if u.op is Ops.MUL) + self.assertIs(invalid_lane_mul.dtype, dtypes.bool) + @unittest.skip("this test isn't valid uops") def test_noop_vectorize_fold(self): d0 = UOp.param(0, dtypes.float, (1,)) diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index 590f1b84cb..b3b04fae3e 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -127,7 +127,7 @@ def do_devectorize(b:UOp): src = [] for idx in itertools.product(*[range(x) for x in b.shape]): idx_c = [UOp.const(None, i) for i in idx] - src.append(b.replace(src=tuple(x.base if x.base.arg is Invalid else x.index(*idx_c) for x in b.src))) + src.append(b.replace(dtype=None, src=tuple(x.base if x.base.arg is Invalid else x.index(*idx_c) for x in b.src))) return UOp.stack(*src).reshape(b.shape) if b.op is not Ops.STORE else UOp.group(*src) def do_stack_wmma(u:UOp):