From 872225e47d784a344a67d990665eb895111d0fb2 Mon Sep 17 00:00:00 2001 From: chenyu Date: Tue, 14 Jul 2026 08:07:00 -0400 Subject: [PATCH] update dtype tests for small dtypes (#17016) --- test/backend/test_dtype.py | 15 ++++++++------- test/helpers.py | 3 ++- 2 files changed, 10 insertions(+), 8 deletions(-) diff --git a/test/backend/test_dtype.py b/test/backend/test_dtype.py index 4b1327c4fc..bbdd395c5b 100644 --- a/test/backend/test_dtype.py +++ b/test/backend/test_dtype.py @@ -25,10 +25,10 @@ def get_available_cast_dtypes(dtype: DType) -> List[DType]: if dtype not in supported_dtypes and dtype not in dtypes.fp8s+(dtypes.half,dtypes.bfloat16): return [] return dts -def _to_torch_storage_type(dtype:DType): - if dtype == dtypes.bfloat16: return torch.float32 - if dtype in dtypes.fp8s: return torch.float32 - return _to_torch_dtype(dtype) +def _to_torch_storage(a:Tensor) -> torch.Tensor: + # tolist() of an fp8 Tensor gives floats, so convert and store in uint8 + if a.dtype in dtypes.fp8s: return torch.tensor([float_to_fp8(x, a.dtype) for x in a.flatten().tolist()], dtype=torch.uint8).reshape(a.shape) + return torch.tensor(a.tolist(), dtype=_to_torch_dtype(a.dtype)) def _test_to_np(a:Tensor, np_dtype, target): if DEBUG >= 2: print(a) @@ -54,7 +54,7 @@ def _test_cast(a:Tensor, target_dtype:DType): if target_dtype in dtypes.fp8s: expected = [truncate[target_dtype](x) for x in expected] _test_op(lambda: a.cast(target_dtype), target_dtype, expected) def _test_bitcast(a:Tensor, target_dtype:DType, target=None): - expected = torch.tensor(a.tolist(), dtype=_to_torch_storage_type(a.dtype)).view(_to_torch_dtype(target_dtype)).tolist() + expected = _to_torch_storage(a).view(_to_torch_dtype(target_dtype)).tolist() if target_dtype in dtypes.fp8s: expected = [fp8_to_float(x, target_dtype) for x in expected] _test_op(lambda: a.bitcast(target_dtype), target_dtype, target or expected) @@ -276,10 +276,11 @@ class TestBitCast(unittest.TestCase): @given(strat.sampled_from(dtype_ints + dtype_floats), strat.sampled_from(dtype_ints + dtype_floats)) def test_shape_change_bitcast(self, dt1, dt2): data = rand_for_dtype(dt1, 32).reshape(2, 2, 8) - expected = torch.tensor(data.tolist(), dtype=_to_torch_storage_type(dt1)).view(_to_torch_dtype(dt2)) + a = Tensor(data, dtype=dt1) + expected = _to_torch_storage(a).view(_to_torch_dtype(dt2)) if dt2 in dtypes.fp8s: expected = torch.tensor([fp8_to_float(x, dt2) for x in expected.view(-1).tolist()]).view_as(expected) - _test_op(lambda: Tensor(data, dtype=dt1).bitcast(dt2), dt2, expected.tolist()) + _test_op(lambda: a.bitcast(dt2), dt2, expected.tolist()) def test_shape_change_bitcast_exceptions(self): with self.assertRaises(RuntimeError): diff --git a/test/helpers.py b/test/helpers.py index 04e28f1900..c4f4c0f73e 100644 --- a/test/helpers.py +++ b/test/helpers.py @@ -6,7 +6,7 @@ from tinygrad import Tensor, dtypes, Device from tinygrad.uop.ops import UOp, Ops, KernelInfo from tinygrad.tensor import _to_np_dtype from tinygrad.codegen import to_program -from tinygrad.dtype import DType +from tinygrad.dtype import DType, truncate from tinygrad.nn.state import get_parameters from tinygrad.helpers import T, Target, DEV from tinygrad.renderer import Renderer @@ -73,6 +73,7 @@ def rand_for_dtype(dt:DType, size:int, allow_subnormal=True): elif dt == dtypes.bool: return np.random.choice([True, False], size=size) ret = np.random.uniform(-10, 10, size=size).astype(_to_np_dtype(dt)) + if dt == dtypes.bfloat16 or dt in dtypes.fp8s: ret = np.array([truncate[dt](x) for x in ret], dtype=ret.dtype) if not allow_subnormal: ret = np.where(np.abs(ret) < min_normal(dt), 0, ret) return ret