mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-11 07:20:33 +02:00
Add dedicated encoder and decoder tests for FP8
This commit is contained in:
+74
-11
@@ -7,6 +7,7 @@
|
||||
#include <assert.h>
|
||||
#include <algorithm>
|
||||
#include <cmath>
|
||||
#include <limits>
|
||||
#include <math.h>
|
||||
#include <stdio.h>
|
||||
#include <string>
|
||||
@@ -155,17 +156,45 @@ static int test_vec_dot_f32(bool verbose) {
|
||||
return num_failed;
|
||||
}
|
||||
|
||||
static float f8_e4m3_to_fp32_ref(uint8_t value) {
|
||||
const uint8_t magnitude = value & 0x7F;
|
||||
if (magnitude == 0x7F) {
|
||||
return NAN;
|
||||
// Positive finite E4M3 values in encoding order, from OCP FP8 rev. 1.1, section 5.1.
|
||||
static const float F8_E4M3_VALUES[] = {
|
||||
0.0f, 0.001953125f, 0.00390625f, 0.005859375f, 0.0078125f, 0.009765625f, 0.01171875f, 0.013671875f,
|
||||
0.015625f, 0.017578125f, 0.01953125f, 0.021484375f, 0.0234375f, 0.025390625f, 0.02734375f, 0.029296875f,
|
||||
0.03125f, 0.03515625f, 0.0390625f, 0.04296875f, 0.046875f, 0.05078125f, 0.0546875f, 0.05859375f,
|
||||
0.0625f, 0.0703125f, 0.078125f, 0.0859375f, 0.09375f, 0.1015625f, 0.109375f, 0.1171875f,
|
||||
0.125f, 0.140625f, 0.15625f, 0.171875f, 0.1875f, 0.203125f, 0.21875f, 0.234375f,
|
||||
0.25f, 0.28125f, 0.3125f, 0.34375f, 0.375f, 0.40625f, 0.4375f, 0.46875f,
|
||||
0.5f, 0.5625f, 0.625f, 0.6875f, 0.75f, 0.8125f, 0.875f, 0.9375f,
|
||||
1.0f, 1.125f, 1.25f, 1.375f, 1.5f, 1.625f, 1.75f, 1.875f,
|
||||
2.0f, 2.25f, 2.5f, 2.75f, 3.0f, 3.25f, 3.5f, 3.75f,
|
||||
4.0f, 4.5f, 5.0f, 5.5f, 6.0f, 6.5f, 7.0f, 7.5f,
|
||||
8.0f, 9.0f, 10.0f, 11.0f, 12.0f, 13.0f, 14.0f, 15.0f,
|
||||
16.0f, 18.0f, 20.0f, 22.0f, 24.0f, 26.0f, 28.0f, 30.0f,
|
||||
32.0f, 36.0f, 40.0f, 44.0f, 48.0f, 52.0f, 56.0f, 60.0f,
|
||||
64.0f, 72.0f, 80.0f, 88.0f, 96.0f, 104.0f, 112.0f, 120.0f,
|
||||
128.0f, 144.0f, 160.0f, 176.0f, 192.0f, 208.0f, 224.0f, 240.0f,
|
||||
256.0f, 288.0f, 320.0f, 352.0f, 384.0f, 416.0f, 448.0f,
|
||||
};
|
||||
|
||||
static uint8_t fp32_to_f8_e4m3_ref(float value) {
|
||||
if (isnan(value)) {
|
||||
return 0x7F;
|
||||
}
|
||||
const int exponent = magnitude >> 3;
|
||||
const int mantissa = magnitude & 0x07;
|
||||
const float result = exponent == 0
|
||||
? ldexpf((float) mantissa, -9)
|
||||
: ldexpf(1.0f + (float) mantissa / 8.0f, exponent - 7);
|
||||
return value & 0x80 ? -result : result;
|
||||
const uint8_t sign = signbit(value) ? 0x80 : 0;
|
||||
const double magnitude = fabs((double) value);
|
||||
if (magnitude >= F8_E4M3_VALUES[126]) {
|
||||
return sign | 126;
|
||||
}
|
||||
int nearest = 0;
|
||||
double distance = magnitude;
|
||||
for (int i = 1; i < 127; ++i) {
|
||||
const double candidate = fabs(magnitude - F8_E4M3_VALUES[i]);
|
||||
if (candidate < distance || (candidate == distance && (i & 1) == 0)) {
|
||||
nearest = i;
|
||||
distance = candidate;
|
||||
}
|
||||
}
|
||||
return sign | (uint8_t) nearest;
|
||||
}
|
||||
|
||||
static int test_f8_e4m3(bool verbose) {
|
||||
@@ -191,8 +220,8 @@ static int test_f8_e4m3(bool verbose) {
|
||||
traits_cpu->from_float(decoded.data(), roundtrip.data(), decoded.size());
|
||||
|
||||
for (int i = 0; i < 256; ++i) {
|
||||
const float expected = f8_e4m3_to_fp32_ref((uint8_t) i);
|
||||
const bool is_nan = (i & 0x7F) == 0x7F;
|
||||
const float expected = is_nan ? NAN : copysignf(F8_E4M3_VALUES[i & 0x7F], i & 0x80 ? -1.0f : 1.0f);
|
||||
const bool decode_failed = is_nan ? !isnan(decoded[i])
|
||||
: decoded[i] != expected || (expected == 0.0f && signbit(decoded[i]) != signbit(expected));
|
||||
const bool encode_failed = !is_nan && roundtrip[i] != encoded[i];
|
||||
@@ -205,6 +234,40 @@ static int test_f8_e4m3(bool verbose) {
|
||||
}
|
||||
}
|
||||
|
||||
std::vector<float> inputs;
|
||||
for (float value : F8_E4M3_VALUES) {
|
||||
inputs.push_back(value);
|
||||
inputs.push_back(-value);
|
||||
}
|
||||
for (int i = 0; i < 126; ++i) {
|
||||
const float midpoint = (F8_E4M3_VALUES[i] + F8_E4M3_VALUES[i + 1]) / 2.0f;
|
||||
for (float value : {nextafterf(midpoint, -INFINITY), midpoint, nextafterf(midpoint, INFINITY)}) {
|
||||
inputs.push_back(value);
|
||||
inputs.push_back(-value);
|
||||
}
|
||||
}
|
||||
for (float value : {std::numeric_limits<float>::denorm_min(), std::numeric_limits<float>::max(),
|
||||
nextafterf(448.0f, INFINITY), nextafterf(464.0f, -INFINITY), 464.0f,
|
||||
nextafterf(464.0f, INFINITY), INFINITY, NAN}) {
|
||||
inputs.push_back(value);
|
||||
inputs.push_back(-value);
|
||||
}
|
||||
|
||||
std::vector<uint8_t> quantized(inputs.size());
|
||||
std::vector<uint8_t> quantized_ref(inputs.size());
|
||||
traits_cpu->from_float(inputs.data(), quantized.data(), inputs.size());
|
||||
traits->from_float_ref(inputs.data(), quantized_ref.data(), inputs.size());
|
||||
for (size_t i = 0; i < inputs.size(); ++i) {
|
||||
const uint8_t expected = fp32_to_f8_e4m3_ref(inputs[i]);
|
||||
if (quantized[i] != expected || quantized_ref[i] != expected) {
|
||||
num_failed++;
|
||||
if (verbose) {
|
||||
printf("f8_e4m3 input %a failed: encoded=0x%02x reference=0x%02x expected=0x%02x\n",
|
||||
inputs[i], quantized[i], quantized_ref[i], expected);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (verbose || num_failed) {
|
||||
printf("f8_e4m3 exhaustive conversion: %s (%d failures)\n", RESULT_STR[num_failed != 0], num_failed);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user