From 7ba386e0418875c333cc01aa8a174c564081969d Mon Sep 17 00:00:00 2001 From: Oliver Simons Date: Wed, 30 Sep 2026 16:59:40 +0200 Subject: [PATCH] Fix and test gguf-py behavior for FP8_E4M3 --- gguf-py/gguf/quants.py | 6 +++--- gguf-py/tests/test_quants.py | 14 ++++++++++++++ 2 files changed, 17 insertions(+), 3 deletions(-) diff --git a/gguf-py/gguf/quants.py b/gguf-py/gguf/quants.py index fe03f52f64..63a137036c 100644 --- a/gguf-py/gguf/quants.py +++ b/gguf-py/gguf/quants.py @@ -768,7 +768,7 @@ class F8_E4M3(__Quant, qtype=GGMLQuantizationType.F8_E4M3): @classmethod def dequantize_blocks(cls, blocks: np.ndarray) -> np.ndarray: bits = blocks.astype(np.uint8) - sign = np.where(bits & 0x80, -1.0, 1.0) + sign = np.where(bits & 0x80, np.float32(-1.0), np.float32(1.0)) magnitude = bits & 0x7F exponent = magnitude >> 3 mantissa = magnitude & 0x07 @@ -777,8 +777,8 @@ class F8_E4M3(__Quant, qtype=GGMLQuantizationType.F8_E4M3): np.ldexp(mantissa.astype(np.float32), -9), np.ldexp(1.0 + mantissa.astype(np.float32) / 8.0, exponent.astype(np.int32) - 7), ) - values = np.where(magnitude == 0x7F, np.nan, values) - return sign * values + values = np.where(magnitude == 0x7F, np.nan, sign * values) + return values.astype(np.float32, copy=False) class IQ2_XXS(__Quant, qtype=GGMLQuantizationType.IQ2_XXS): diff --git a/gguf-py/tests/test_quants.py b/gguf-py/tests/test_quants.py index 9aa7c4ae2a..53a5cd2ffe 100755 --- a/gguf-py/tests/test_quants.py +++ b/gguf-py/tests/test_quants.py @@ -69,6 +69,7 @@ class GGMLQuants: "tq1_0", "tq2_0", "mxfp4", "nvfp4", + "f8_e4m3", "iq2_xxs", "iq2_xs", "iq2_s", "iq3_xxs", "iq3_s", "iq1_s", "iq1_m", "iq4_nl", "iq4_xs", ): @@ -180,6 +181,19 @@ def do_test(libggml_path: Path, quick: bool = False, user_type: GGMLQuantization logger.info(f"Testing {qtype.name}") + if qtype == GGMLQuantizationType.F8_E4M3: + encodings = np.arange(256, dtype=np.uint8) + pydq = gguf.dequantize(encodings, qtype) + ggdq = ggml_quants.dequantize(encodings, qtype) + assert pydq.dtype == np.float32 + nan_mask = (encodings & 0x7F) == 0x7F + # Both 0x7F and 0xFF are NaNs; check classification, not FP32 NaN sign or payload. + np.testing.assert_array_equal(np.isnan(pydq), nan_mask) + np.testing.assert_array_equal(np.isnan(ggdq), nan_mask) + # Compare finite values by bits to also check signed zeros. + np.testing.assert_array_equal(pydq[~nan_mask].view(np.uint32), ggdq[~nan_mask].view(np.uint32)) + logger.info("All 256 FP8 E4M3 encodings match C") + rc = r.copy(order="C") pyq = None