File size: 3,644 Bytes
bbb6388 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 | // src/kernels/bf16_bits_test.cpp - `bf16_bits.hpp` against ggml's rule, over EVERY f32 bit pattern (CPU only).
//
// the rounding add `i + 0x7FFF + bit16` carries out of the mantissa for a NaN, so 0x7FFFFFFF came back as
// -0 and 0x7F800001 as +inf - an upstream NaN was masked instead of propagated. `ggml_compute_fp32_to_bf16`
// (ggml-impl.h at the pinned llama.cpp) tests the NaN first and returns `(i >> 16) | 64`. Three checks, all
// exhaustive (2^32 inputs, a few seconds):
//
// 1. the header equals ggml's function, transcribed below, on every input;
// 2. on every NON-NaN input it equals the previous rounding, so the change cannot move a finite value or an
// infinity by a single bit;
// 3. every NaN input gives a quiet NaN of the same sign.
//
// And the round trip over all 65,536 bf16 patterns: a non-NaN bf16 survives f32 and back unchanged.
#include "strata/kernels/bf16_bits.hpp"
#include <cstdint>
#include <cstdio>
#include <cstring>
namespace {
/// ggml_compute_fp32_to_bf16, ggml/src/ggml-impl.h (MIT, the ggml authors), on raw bits.
uint16_t ggml_bf16(uint32_t u) {
if ((u & 0x7fffffffu) > 0x7f800000u) return (uint16_t) ((u >> 16) | 64);
return (uint16_t) ((u + (0x7fffu + ((u >> 16) & 1u))) >> 16);
}
/// The rule before this fix, kept only to prove the change is confined to NaN inputs.
uint16_t previous_bf16(uint32_t u) {
u = (u + ((u >> 16) & 1u) + 0x7FFFu) & 0xFFFF0000u;
return (uint16_t) (u >> 16);
}
bool is_nan_bits(uint32_t u) { return (u & 0x7fffffffu) > 0x7f800000u; }
} // namespace
int main() {
unsigned long long vs_ggml = 0, vs_previous = 0, nan_lost = 0, nan_inputs = 0, previous_masked = 0;
uint32_t first_bad = 0;
bool have_bad = false;
for (uint64_t k = 0; k <= 0xFFFFFFFFull; ++k) {
const uint32_t u = (uint32_t) k;
float f;
std::memcpy(&f, &u, 4);
const uint16_t got = strata::kernels::bf16_from_f32(f);
if (got != ggml_bf16(u)) {
++vs_ggml;
if (!have_bad) { first_bad = u; have_bad = true; }
}
if (is_nan_bits(u)) {
++nan_inputs;
const uint32_t back = (uint32_t) got << 16;
// a quiet NaN (bit 22 of the f32) with the input's sign
if (!is_nan_bits(back) || (back & 0x00400000u) == 0 || (back >> 31) != (u >> 31)) ++nan_lost;
if (!is_nan_bits((uint32_t) previous_bf16(u) << 16)) ++previous_masked;
} else if (got != previous_bf16(u)) {
++vs_previous;
if (!have_bad) { first_bad = u; have_bad = true; }
}
}
std::printf(" f32 -> bf16 over 2^32 inputs: %llu differ from ggml, %llu non-NaN differ from the previous rule\n",
vs_ggml, vs_previous);
std::printf(" NaN inputs: %llu, not a quiet same-sign NaN now: %llu (the previous rule masked %llu of them)\n",
nan_inputs, nan_lost, previous_masked);
if (have_bad) std::printf(" first mismatch at input 0x%08x\n", first_bad);
unsigned long long trip_bad = 0;
for (uint32_t h = 0; h <= 0xFFFFu; ++h) {
const float f = strata::kernels::f32_from_bf16((uint16_t) h);
const uint16_t back = strata::kernels::bf16_from_f32(f);
const bool nan = is_nan_bits(h << 16);
if (back != (nan ? (uint16_t) (h | 64u) : (uint16_t) h)) ++trip_bad;
}
std::printf(" bf16 -> f32 -> bf16 over 65536 patterns: %llu wrong\n", trip_bad);
const bool ok = vs_ggml == 0 && vs_previous == 0 && nan_lost == 0 && trip_bad == 0 && previous_masked > 0;
std::printf("\nbf16_bits_test %s\n", ok ? "OK" : "FAILED");
return ok ? 0 : 1;
}
|