diff --git a/test/null/test_dtype_spec.py b/test/null/test_dtype_spec.py index 8b8e32bb76..ac04ae5ca6 100644 --- a/test/null/test_dtype_spec.py +++ b/test/null/test_dtype_spec.py @@ -58,8 +58,8 @@ class TestHelpers(unittest.TestCase): def test_from_py(self): assert dtypes.from_py(True) == dtypes.bool assert dtypes.from_py(Invalid) == dtypes.bool - assert dtypes.from_py(2) == dtypes.default_int - assert dtypes.from_py(3.0) == dtypes.default_float + assert dtypes.from_py(2) == dtypes.weakint + assert dtypes.from_py(3.0) == dtypes.weakfloat assert dtypes.from_py([]) == dtypes.default_float assert dtypes.from_py(()) == dtypes.default_float assert dtypes.from_py([True]) == dtypes.bool @@ -313,15 +313,15 @@ class TestAutoCastType(unittest.TestCase): @given(strat.sampled_from(core_dtypes)) def test_broadcast_scalar(self, dt): - assert (Tensor.ones(4, 4, dtype=dt) + 2.3).dtype == (dt if dtypes.is_float(dt) else dtypes.default_float) - assert (Tensor.ones(4, 4, dtype=dt) + 2).dtype == (dt if dtypes.is_float(dt) or dtypes.is_int(dt) else dtypes.default_int) + assert (Tensor.ones(4, 4, dtype=dt) + 2.3).dtype == (dt if dtypes.is_float(dt) else dtypes.weakfloat) + assert (Tensor.ones(4, 4, dtype=dt) + 2).dtype == (dt if dtypes.is_float(dt) or dtypes.is_int(dt) else dtypes.weakint) assert (Tensor.ones(4, 4, dtype=dt) + True).dtype == dt @given(strat.sampled_from(core_dtypes)) def test_pad_scalar(self, dt): t = Tensor.ones(4, dtype=dt) - assert t.pad(((1, 1),), value=2.3).dtype == (dt if dtypes.is_float(dt) else dtypes.default_float) - assert t.pad(((1, 1),), value=2).dtype == (dt if dtypes.is_float(dt) or dtypes.is_int(dt) else dtypes.default_int) + assert t.pad(((1, 1),), value=2.3).dtype == (dt if dtypes.is_float(dt) else dtypes.weakfloat) + assert t.pad(((1, 1),), value=2).dtype == (dt if dtypes.is_float(dt) or dtypes.is_int(dt) else dtypes.weakint) assert t.pad(((1, 1),), value=True).dtype == dt @given(strat.sampled_from(core_dtypes)) @@ -420,16 +420,16 @@ class TestAutoCastType(unittest.TestCase): @given(strat.sampled_from(core_dtypes)) def test_where_one_scalar(self, dt): t = Tensor(2, dtype=dt) - self.check_where_alternate_input_other(t, 3.2, (dt if dtypes.is_float(dt) else dtypes.default_float)) - self.check_where_alternate_input_other(t, 3, (dt if dtypes.is_float(dt) or dtypes.is_int(dt) else dtypes.default_int)) + self.check_where_alternate_input_other(t, 3.2, (dt if dtypes.is_float(dt) else dtypes.weakfloat)) + self.check_where_alternate_input_other(t, 3, (dt if dtypes.is_float(dt) or dtypes.is_int(dt) else dtypes.weakint)) self.check_where_alternate_input_other(t, True, dt) def test_where_two_scalars(self): - self.check_where_alternate_input_other(3.1, 3.2, dtypes.default_float) - self.check_where_alternate_input_other(3.1, 3, dtypes.default_float) - self.check_where_alternate_input_other(3.1, True, dtypes.default_float) - self.check_where_alternate_input_other(3, 2, dtypes.default_int) - self.check_where_alternate_input_other(3, True, dtypes.default_int) + self.check_where_alternate_input_other(3.1, 3.2, dtypes.weakfloat) + self.check_where_alternate_input_other(3.1, 3, dtypes.weakfloat) + self.check_where_alternate_input_other(3.1, True, dtypes.weakfloat) + self.check_where_alternate_input_other(3, 2, dtypes.weakint) + self.check_where_alternate_input_other(3, True, dtypes.weakint) def test_where_non_bool_cond_raises(self): with self.assertRaises(RuntimeError): Tensor([1, 0, 2]).where(1, 0) @@ -441,8 +441,8 @@ class TestAutoCastType(unittest.TestCase): @given(strat.sampled_from(core_dtypes)) def test_maximum_const(self, dt): - assert Tensor([1, 2], dtype=dt).maximum(3.1).dtype == (dt if dtypes.is_float(dt) else dtypes.default_float) - assert Tensor([1, 2], dtype=dt).maximum(3).dtype == (dt if dtypes.is_float(dt) or dtypes.is_int(dt) else dtypes.default_int) + assert Tensor([1, 2], dtype=dt).maximum(3.1).dtype == (dt if dtypes.is_float(dt) else dtypes.weakfloat) + assert Tensor([1, 2], dtype=dt).maximum(3).dtype == (dt if dtypes.is_float(dt) or dtypes.is_int(dt) else dtypes.weakint) assert Tensor([1, 2], dtype=dt).maximum(True).dtype == dt def test_div(self): @@ -453,7 +453,7 @@ class TestAutoCastType(unittest.TestCase): def test_div_const(self): assert (Tensor([1, 2], dtype=dtypes.int32) / 2).dtype == dtypes.default_float - assert (Tensor([1, 2], dtype=dtypes.int32) / 2.0).dtype == dtypes.default_float + assert (Tensor([1, 2], dtype=dtypes.int32) / 2.0).dtype == dtypes.weakfloat assert (Tensor([1, 2], dtype=dtypes.float16) / 2).dtype == dtypes.float16 assert (Tensor([1, 2], dtype=dtypes.float16) / 2.0).dtype == dtypes.float16 diff --git a/test/null/test_tensor.py b/test/null/test_tensor.py index 9090eccc5f..4156b37177 100644 --- a/test/null/test_tensor.py +++ b/test/null/test_tensor.py @@ -181,7 +181,7 @@ class TestTensorPad(unittest.TestCase): t = Tensor.arange(9).reshape(1, 1, 3, 3) self.assertEqual(t.dtype, dtypes.int) r = t.pad((1, 2, 0, -1), value=-float('inf')) - self.assertEqual(r.dtype, dtypes.float) + self.assertEqual(r.dtype, dtypes.weakfloat) self.assertEqual(r.shape, (1, 1, 2, 6)) class TestTensorDeviceMismatch(unittest.TestCase): diff --git a/test/unit/test_dtype_spec.py b/test/unit/test_dtype_spec.py index 93c15ea372..05bd2df703 100644 --- a/test/unit/test_dtype_spec.py +++ b/test/unit/test_dtype_spec.py @@ -1,6 +1,6 @@ import unittest, math, subprocess from tinygrad.tensor import Tensor -from tinygrad.dtype import dtypes, DType, DTYPES_DICT +from tinygrad.dtype import dtypes, DType, DTYPES_DICT, strong_dtype from tinygrad.device import Device from tinygrad.helpers import getenv, DEBUG, EMULATED_DTYPES from test.helpers import slow @@ -25,6 +25,8 @@ def _assert_eq(tensor:Tensor, target_dtype:DType, target, tol_target_dtype:float if DEBUG >= 2: print(tensor.numpy()) try: assert tensor.dtype == target_dtype + # weak values read back at their default. + target_dtype = strong_dtype(target_dtype) # denormals are zero if target_dtype in dtypes.floats and (target_dtype not in supported_dtypes or target_dtype in EMULATED_DTYPES.tolist(dtypes)): fe, fm = dtypes.finfo(target_dtype) @@ -82,8 +84,8 @@ class TestTypeSpec(unittest.TestCase): dtypes.default_int, dtypes.default_float = default_int, default_float _assert_eq(Tensor(True), dtypes.bool, True) _assert_eq(Tensor(None), dtypes.default_float, []) - _assert_eq(Tensor(2), dtypes.default_int, 2) - _assert_eq(Tensor(2.34), dtypes.default_float, 2.34) + _assert_eq(Tensor(2), dtypes.weakint, 2) + _assert_eq(Tensor(2.34), dtypes.weakfloat, 2.34) _assert_eq(Tensor([]), dtypes.default_float, []) _assert_eq(Tensor([1]), dtypes.default_int, [1]) # list elements are python scalars; a numpy scalar in a list has no inferred dtype (use np.array or state a dtype) diff --git a/test/unit/test_dtype_weak.py b/test/unit/test_dtype_weak.py index eb18d12012..de7840c0ef 100644 --- a/test/unit/test_dtype_weak.py +++ b/test/unit/test_dtype_weak.py @@ -14,10 +14,11 @@ class TestWeakPromotion(unittest.TestCase): with self.assertRaises(ValueError): Tensor.const(dtypes.weakfloat, 1.0).rand_like() with self.assertRaises(ValueError): Tensor.const(dtypes.weakfloat, 1.0).randn_like() - def test_sum_stays_weak(self): - for weak, value in ((dtypes.weakfloat, 1.0),): - self.assertEqual(Tensor.const(weak, value).expand(3).sum().dtype, weak) - self.assertEqual((Tensor.const(dtypes.weakfloat, 1.0).expand(3).sum() + Tensor([1], dtype=dtypes.float16)).dtype, dtypes.float16) + def test_reduce_strips_weakness(self): + for weak, value, strong in ((dtypes.weakint, 1, dtypes.default_int), (dtypes.weakfloat, 1.0, dtypes.default_float)): + t = Tensor.const(weak, value).expand(3) + for out in (t.sum(), t.max(), t.prod(), t.cumsum(0), t.cummax(0)[0]): self.assertEqual(out.dtype, strong) + self.assertEqual((Tensor.const(dtypes.weakfloat, 1.0).expand(3).sum() + Tensor([1], dtype=dtypes.float16)).dtype, dtypes.float32) def test_materialize_at_default_dtype(self): for weak, value, strong in ((dtypes.weakfloat, 0.5, dtypes.default_float),): @@ -67,7 +68,10 @@ class TestWeakPromotion(unittest.TestCase): self.assertEqual((t_f32 + t_f16).dtype, dtypes.float32) self.assertEqual(Tensor([2], dtype=dtypes.uint8).pad(((1, 1),), value=1).dtype, dtypes.uint8) - @unittest.expectedFailure # TODO: dot of a weak const tensor defers to the other operand once python scalars are weak consts + def test_concrete_pair_promotes_weak(self): + out = Tensor([-1], dtype=dtypes.int64, device="CPU") + Tensor([3], dtype=dtypes.uint64, device="CPU") + Tensor(0.5) + self.assertEqual((out.dtype, out.tolist()), (dtypes.weakfloat, [2.5])) + def test_dot_defers_weak(self): weak = Tensor([True, False]).where(Tensor(1), 2) self.assertEqual(weak.dot(Tensor([1, 1], dtype=dtypes.int8)).dtype, dtypes.int8) @@ -102,7 +106,6 @@ class TestWeakPromotion(unittest.TestCase): x32 = Tensor.full((1,), 0.0, dtype=dtypes.float32, device="CPU") self.assertEqual((x32 + value).item(), 1.0) - @unittest.expectedFailure # TODO: exp/cos/sigmoid of a weak const stay weak instead of casting to a concrete float def test_weak_transcendentals(self): t_f16 = Tensor([1], dtype=dtypes.float16) for out in (Tensor(2).exp(), Tensor(2).cos(), Tensor(2).sigmoid()): diff --git a/tinygrad/codegen/decomp/dtype.py b/tinygrad/codegen/decomp/dtype.py index 39b21c9070..5c9c7c1200 100644 --- a/tinygrad/codegen/decomp/dtype.py +++ b/tinygrad/codegen/decomp/dtype.py @@ -24,7 +24,10 @@ def l2i(op: Ops, dt: DType, *uops:UOp): match op: case Ops.NEG: return l2i(Ops.SUB, dt, zero, zero, *uops) case Ops.CAST if dt in (dtypes.long, dtypes.ulong) and uops[0].dtype not in dtypes.floats: - return uops[0].cast(l2i_dt[dt]), (uops[0] < 0).where(UOp.const(l2i_dt[dt], -1), UOp.const(l2i_dt[dt], 0)) + # the high word is the sign extension; bool has no sign, test the already-cast low word instead (bool < 0 would promote to weakint) + x, lo = uops[0], uops[0].cast(l2i_dt[dt]) + sign = lo if x.dtype is dtypes.bool else x + return lo, (sign < sign.const_like(0)).where(lo.const_like(-1), lo.const_like(0)) case Ops.CAST if dt in (dtypes.long, dtypes.ulong): return (lo:=uops[0].cast(l2i_dt[dt])), (uops[0] / 2**32).cast(l2i_dt[dt]) - ((uops[0] < 0) & lo.ne(0)).cast(l2i_dt[dt]) case Ops.CAST if dt in dtypes.floats: diff --git a/tinygrad/dtype.py b/tinygrad/dtype.py index 60fa045490..eddc38a4f2 100644 --- a/tinygrad/dtype.py +++ b/tinygrad/dtype.py @@ -99,8 +99,8 @@ class dtypes: def from_py(x) -> DType: # NOTE: isinstance(True, int) is True, so bool must be checked before int if isinstance(x, (bool, InvalidType)): return dtypes.bool - if isinstance(x, float): return dtypes.default_float - if isinstance(x, int): return dtypes.default_int + if isinstance(x, float): return dtypes.weakfloat + if isinstance(x, int): return dtypes.weakint # put this in the last is faster because there are more items than lists/tuples to check if isinstance(x, (list, tuple)): return strong_dtype(max(dtypes.from_py(xi) for xi in x)) if x else dtypes.default_float raise RuntimeError(f"Could not infer dtype of {x} with type {type(x)}") @@ -208,7 +208,6 @@ def can_lossless_cast(dt0:DType, dt1:DType) -> bool: def sum_acc_dtype(dt:DType): # default acc dtype for sum - if dt in dtypes.weaks: return dt if dtypes.is_unsigned(dt): return least_upper_dtype(dt, dtypes.uint) if dtypes.is_int(dt) or dt == dtypes.bool: return least_upper_dtype(dt, dtypes.int) return least_upper_dtype(dt, to_dtype(getenv("SUM_DTYPE", "float32"))) @@ -303,6 +302,7 @@ def _from_np_dtype(npdtype:'np.dtype') -> DType: # type: ignore [name-defined] # @functools.cache def _to_torch_dtype(dtype:DType) -> 'torch.dtype'|None: # type: ignore [name-defined] # noqa: F821 import numpy as np, torch + dtype = strong_dtype(dtype) if dtype == dtypes.uint64: return torch.uint64 if dtype == dtypes.bfloat16: return torch.bfloat16 if dtype in dtypes.fp8s: return torch.uint8 diff --git a/tinygrad/mixin/reduce.py b/tinygrad/mixin/reduce.py index 6f5285968b..544fa56ce3 100644 --- a/tinygrad/mixin/reduce.py +++ b/tinygrad/mixin/reduce.py @@ -1,6 +1,6 @@ from typing import Self, Sequence from tinygrad.uop import Ops -from tinygrad.dtype import DTypeLike, dtypes, sum_acc_dtype, to_dtype +from tinygrad.dtype import DTypeLike, dtypes, strong_dtype, sum_acc_dtype, to_dtype from tinygrad.helpers import make_tuple from tinygrad.mixin.dtype import DTypeMixin from tinygrad.mixin.movement import MovementMixin @@ -11,6 +11,7 @@ class ReduceMixin(DTypeMixin, MovementMixin): raise NotImplementedError def _reduce(self, op:Ops, axis:int|Sequence[int]|None=None, keepdim=False) -> Self: + if self.dtype in dtypes.weaks: self = self.cast(strong_dtype(self.dtype)) axis = tuple(self._resolve_dim(x) for x in (range(self.ndim) if axis is None else make_tuple(axis, 1))) if self.ndim == 0: axis = () ret = self._rop(op, axis) diff --git a/tinygrad/renderer/isa/x86.py b/tinygrad/renderer/isa/x86.py index a19a2fbdc5..6ded880f02 100644 --- a/tinygrad/renderer/isa/x86.py +++ b/tinygrad/renderer/isa/x86.py @@ -193,8 +193,9 @@ pre_isel_matcher = PatternMatcher([ (UPat((Ops.INDEX, Ops.SHRINK), name="addr").store(UPat.var("val"), UPat.var("gate")), gated_store), # TODO: remove this once we allow all flag producing ops in cmove # if gate in scalar int cmove is not a comparison need to add one to set the flag + # NOTE: the 0 is int so the bool gate zero-extends and compares as int (a byte compare renders different kernels) (UPat.var("m", dtypes.bool).where(UPat.var("a"), UPat.var("b")), - lambda m,a,b: m.ne(0).where(a,b) if m.op not in GroupOp.Comparison else None), + lambda m,a,b: m.ne(UOp.const(dtypes.int, 0)).where(a,b) if m.op not in GroupOp.Comparison else None), ]) # ***** X86 registers ***** diff --git a/tinygrad/renderer/wgsl.py b/tinygrad/renderer/wgsl.py index 20e3b7b59e..d8620aed45 100644 --- a/tinygrad/renderer/wgsl.py +++ b/tinygrad/renderer/wgsl.py @@ -13,6 +13,8 @@ def sign_extend(val:UOp, sext_am:int): def packed_store(bidx:UOp, var:UOp, gate:UOp|None=None): elems, mask = 4//var.dtype.itemsize, _mask(var.dtype) shift_am, div_idx = (bidx.src[1].cast(dtypes.uint32) % elems) * (8*var.dtype.itemsize), bidx.src[1] // elems + # bool does its mask math at int32: renderer rewrites run after weak dtypes are lowered, and bool & 0xFF would create a weakint const + if var.dtype == dtypes.bool: var = var.cast(dtypes.int32) new_v, wmask = (var & mask).cast(dtypes.uint32) << shift_am, ((mask << shift_am) ^ 0xFFFFFFFF).cast(dtypes.uint32) idx = UOp(Ops.INDEX, src=(bidx.src[0], div_idx)) buf = UOp.load(idx, *((UOp.const(dtypes.uint32, 0), gate) if gate is not None else ()), dtype=dtypes.uint32) diff --git a/tinygrad/schedule/rangeify.py b/tinygrad/schedule/rangeify.py index bb88fb094c..3f8311dcb1 100644 --- a/tinygrad/schedule/rangeify.py +++ b/tinygrad/schedule/rangeify.py @@ -1,7 +1,7 @@ from dataclasses import dataclass, field, replace from typing import cast import itertools -from tinygrad.dtype import dtypes, AddrSpace, Invalid, to_dtype +from tinygrad.dtype import dtypes, AddrSpace, Invalid, to_dtype, strong_dtype from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, resolve, GroupOp, KernelInfo, ParamArg, shape_to_shape_arg from tinygrad.uop.ops import graph_rewrite, sint, AxisType, BottomUpGate, profile_matches, identity_element from tinygrad.uop.symbolic import symbolic @@ -359,6 +359,7 @@ pm_limit_bufs = PatternMatcher([(UPat(set.union(GroupOp.Binary, GroupOp.Ternary) def bufferize_to_store(ctx:itertools.count, x:UOp, idx:UOp, allow_locals=True): size = prod(x.shape) + dtype = strong_dtype(x.dtype) # a BUFFER is never weak: store at the concrete dtype, the .cast(x.dtype) on the result keeps readers unchanged rngs = sorted(idx.ranges, key=lambda x: x.arg) assert size > 0 and isinstance(size, int), f"no zero sized or symbolic sized buffers {size}" @@ -379,15 +380,15 @@ def bufferize_to_store(ctx:itertools.count, x:UOp, idx:UOp, allow_locals=True): # NOTE: the local BUFFER needs to be disambiguated here if x.arg.addrspace == AddrSpace.GLOBAL: - buf = UOp(Ops.BUFFER, src=(shape_to_shape_arg((size,)),), arg=ParamArg(next(ctx), x.dtype, device=x.arg.device, addrspace=AddrSpace.GLOBAL)) - do_store = buf.index(idx).store(x.src[0]).end(*rngs) - return buf.after(do_store) + buf = UOp(Ops.BUFFER, src=(shape_to_shape_arg((size,)),), arg=ParamArg(next(ctx), dtype, device=x.arg.device, addrspace=AddrSpace.GLOBAL)) + do_store = buf.index(idx).store(x.src[0].cast(dtype)).end(*rngs) + return buf.after(do_store).cast(x.dtype) if allow_locals: # handle locals - buf = UOp.placeholder((size,), x.dtype, next(ctx), AddrSpace.LOCAL) - do_store = buf.index(idx).store(x.src[0]).end(*rngs) - return buf.after(do_store.barrier()) + buf = UOp.placeholder((size,), dtype, next(ctx), AddrSpace.LOCAL) + do_store = buf.index(idx).store(x.src[0].cast(dtype)).end(*rngs) + return buf.after(do_store.barrier()).cast(x.dtype) # collapse any BUFFERIZE to single input BUFFERIZE def flatten_bufferize(x:UOp): @@ -412,6 +413,11 @@ def remove_noop_afters(x:UOp) -> UOp|None: pm_add_buffers = pm_mops+pm_flatten_bufferize+PatternMatcher([ (UPat(Ops.STAGE, src=(UPat(), UPat(name="idx")), name="x"), lambda ctx,x,idx: bufferize_to_store(ctx, x, idx, allow_locals=False)), + # INDEX of a buffer through the weak cast added above: index the buffer directly and cast the loaded value instead. + # this must run in the same rewrite that adds the cast, or the expander expands the whole casted buffer into one big VECTORIZE + (UPat(Ops.INDEX, src=(UPat(Ops.CAST, dtype=dtypes.weaks, src=(UPat.var("buf"),)),), allow_any_len=True, name="u"), + lambda u,buf: u.replace(dtype=None, src=(buf,)+u.src[1:]).cast(u.dtype)), + # move RESHAPEs through MSELECT/MSTACK (UPat((Ops.MSELECT, Ops.MSTACK), src=UPat(Ops.RESHAPE), name="m"), lambda m: m.replace(src=tuple([x.src[0].base for x in m.src])).reshape(m.shape)), diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index c438917631..e82b0b263d 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -1728,6 +1728,10 @@ def lower_weak_node(u:UOp) -> UOp|None: return u.replace(dtype=None, src=src[:start]+tuple(s.cast(dt) for s in src[start:])).cast(u.dtype) pm_lower_weak = PatternMatcher([ (UPat(Ops.CONST, dtype=dtypes.weaks, name="u"), lambda u: u.replace(dtype=select_dtype(u)).cast(u.dtype)), + # two stacked weak casts are a weakint value used as weakfloat (or vice versa): resolve the inner one at the outer kind's default. + # a SINGLE weak cast is never rewritten here, each consumer absorbs it on its own edge (see lower_weak_srcs) + (UPat(Ops.CAST, dtype=dtypes.weaks, src=(UPat(Ops.CAST, dtype=dtypes.weaks, src=(UPat.var("x"),)),), name="u"), + lambda u,x: x.cast(select_dtype(u)).cast(u.dtype) if x.dtype not in dtypes.weaks else None), # Binary can widen from the bounds, all other nodes derive from the lowered sources. # a weakfloat Unary (sin/exp2/...) must resolve here, before the transcendental decomposition (UPat(GroupOp.Binary|GroupOp.Unary|{Ops.WHERE, Ops.RANGE, Ops.STACK}, name="u"), lower_weak_node),