diff --git a/test/unit/test_dtype_weak.py b/test/unit/test_dtype_weak.py index f85a64dc12..10dd40c0d0 100644 --- a/test/unit/test_dtype_weak.py +++ b/test/unit/test_dtype_weak.py @@ -79,11 +79,13 @@ class TestWeakPromotion(unittest.TestCase): self.assertEqual((v & 3).dtype, dtypes.weakint) with self.assertRaises(RuntimeError): Tensor.const(dtypes.weakfloat, 1.0) << Tensor.const(dtypes.weakfloat, 1.0) with self.assertRaises(RuntimeError): UOp.const(dtypes.int32, 1).alu(Ops.SHL, UOp.const(dtypes.float64, 1)) - # float bitwise/shift builds, the spec rejects it + for op in (Ops.SHL, Ops.SHR): + with self.assertRaises(RuntimeError): + UOp.const(dtypes.float32, 1).alu(op, UOp.const(dtypes.int32, 1)) + # float bitwise builds, the spec rejects it with Context(SPEC=1): f32, wf = UOp.const(dtypes.float32, 1.0), UOp.const(dtypes.weakfloat, 1.0) - for bad in (f32.alu(Ops.AND, f32), f32.alu(Ops.SHL, UOp.const(dtypes.int32, 1)), - UOp(Ops.AND, dtypes.float32, (f32, f32)), UOp(Ops.AND, dtypes.int32, (wf, wf))): + for bad in (f32.alu(Ops.AND, f32), UOp(Ops.AND, dtypes.float32, (f32, f32)), UOp(Ops.AND, dtypes.int32, (wf, wf))): with self.assertRaises(RuntimeError): type_verify([bad], spec_shared) def test_integer_values(self): @@ -144,8 +146,9 @@ class TestWeakMaterializationEntries(unittest.TestCase): self.assertEqual(weak_val().tolist(), [value]) self.assertEqual(weak_val().cast(strong).realize().uop.buffer.dtype, strong) self.assertEqual(weak_val().contiguous().dtype, weak) # no layout to fix, stays weak - self.assertEqual(weak_val().clone().dtype, strong) # storage commits at the kind default - for entry in (lambda t: t.realize(), lambda t: t.to("CPU:1").realize(), lambda t: t.as_param(0)): + self.assertEqual(weak_val().realize().dtype, weak) # no width to store, stays weak + self.assertEqual(weak_val().clone().dtype, strong) # storage commits at the default + for entry in (lambda t: t.to("CPU:1").realize(), lambda t: t.as_param(0)): with self.assertRaises(RuntimeError): entry(weak_val()) def test_empty_reads_commit(self): diff --git a/tinygrad/tensor.py b/tinygrad/tensor.py index 2cb78de949..b905717b4a 100644 --- a/tinygrad/tensor.py +++ b/tinygrad/tensor.py @@ -189,7 +189,8 @@ class Tensor(RandMixin): @disable_gc() def realize(self, *lst:Tensor, do_update_stats=True) -> Tensor: """Triggers the computation needed to create these Tensor(s).""" - if len(to_realize:=[x for x in (self,)+lst if x.uop.device is not None and not x.uop.has_buffer_identity()]): + to_realize = [x for x in (self,)+lst if x.uop.device is not None and not x.uop.has_buffer_identity() and x.dtype not in dtypes.weaks] + if len(to_realize): run_linear(*Tensor.linear_with_vars(*to_realize), update_stats=do_update_stats) return self diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index e714b59ca5..6a006b8bf1 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -156,7 +156,7 @@ def dtype_from_uop(op:Ops, src:tuple[UOp,...], arg:Any) -> DType|None: case Ops.GETADDR: return dtypes.uint64 case Ops.SHL | Ops.SHR: - if not dtypes.is_int(src[1].dtype): raise RuntimeError(f"shift distance must be int, got {src[1].dtype}") + if not all(dtypes.is_int(x.dtype) for x in src): raise RuntimeError(f"shift operands must be int, got {[x.dtype for x in src]}") return promo_dtype(src) case Ops.BUFFER | Ops.PARAM: assert isinstance(arg, ParamArg), "BUFFER/PARAM must have ParamArg" @@ -870,6 +870,7 @@ class UOp(RandMixin, metaclass=UOpMetaClass): def has_buffer_identity(self, after_ok=False): """Check if this UOp has a concrete buffer identity in the graph (RESHAPE/MULTI -> BUFFER chain).""" + # TODO: this is confusing because UOp.variable('v', 0, 1, dtypes.weakfloat) is True for jit to work, but it doesn't have a buffer if self.op in {Ops.RESHAPE, Ops.MULTI, Ops.MSELECT}: return self.src[0].has_buffer_identity(after_ok) if after_ok and self.op == Ops.AFTER: return self.src[0].has_buffer_identity(after_ok) return self.op in {Ops.BUFFER, Ops.SLICE, Ops.PARAM}