From ac3449b0c87fbf80237dc45d40d053b40e32cf39 Mon Sep 17 00:00:00 2001 From: chenyu Date: Mon, 25 Aug 2025 19:03:41 -0400 Subject: [PATCH] truncate_fp16 cleanup (#11838) native `@` is default --- test/unit/test_dtype_spec.py | 7 ++++++- tinygrad/dtype.py | 2 +- 2 files changed, 7 insertions(+), 2 deletions(-) diff --git a/test/unit/test_dtype_spec.py b/test/unit/test_dtype_spec.py index 0e78af19d4..0353acc230 100644 --- a/test/unit/test_dtype_spec.py +++ b/test/unit/test_dtype_spec.py @@ -102,13 +102,18 @@ class TestHelpers(unittest.TestCase): self.assertEqual(truncate_fp16(65504), 65504) self.assertEqual(truncate_fp16(65519.999), 65504) self.assertEqual(truncate_fp16(65520), math.inf) + self.assertEqual(truncate_fp16(1e-8), 0.0) + self.assertEqual(truncate_fp16(-65504), -65504) + self.assertEqual(truncate_fp16(-65519.999), -65504) + self.assertEqual(truncate_fp16(-65520), -math.inf) + self.assertTrue(math.isnan(truncate_fp16(math.nan))) def test_truncate_bf16(self): self.assertEqual(truncate_bf16(1), 1) + # TODO: rounding, torch bfloat 1.1 gives 1.1015625 instead of 1.09375 self.assertAlmostEqual(truncate_bf16(1.1), 1.09375, places=7) for a in [1234, 23456, -777.777]: self.assertEqual(truncate_bf16(a), torch.tensor([a], dtype=torch.bfloat16).item()) - # TODO: torch bfloat 1.1 gives 1.1015625 instead of 1.09375 max_bf16 = torch.finfo(torch.bfloat16).max self.assertEqual(truncate_bf16(max_bf16), max_bf16) self.assertEqual(truncate_bf16(min_bf16:=-max_bf16), min_bf16) diff --git a/tinygrad/dtype.py b/tinygrad/dtype.py index 2fb651763f..5074179c1f 100644 --- a/tinygrad/dtype.py +++ b/tinygrad/dtype.py @@ -215,7 +215,7 @@ def sum_acc_dtype(dt:DType): return least_upper_dtype(dt, to_dtype(getenv("SUM_DTYPE", "float32"))) def truncate_fp16(x): - try: return struct.unpack("@e", struct.pack("@e", float(x)))[0] + try: return struct.unpack('e', struct.pack('e', float(x)))[0] except OverflowError: return math.copysign(math.inf, x) def truncate_bf16(x):