mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-02 02:47:26 -05:00
Fix and test gguf-py behavior for FP8_E4M3
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user