#ifndef VT_FP8_KV_H_
#define VT_FP8_KV_H_
#include <cmath>
#include <cstdint>
#include <limits>
namespace vt {
enum class Fp8KVCacheDataType : uint8_t {
kAuto = 0, kFp8E4M3, kFp8E5M2, };
inline float F8E4M3ToF32(uint8_t byte) {
const uint32_t sign = static_cast<uint32_t>(byte >> 7) & 0x1U;
const uint32_t exp = static_cast<uint32_t>(byte >> 3) & 0xFU;
const uint32_t mant = static_cast<uint32_t>(byte) & 0x7U;
const float sm = sign ? -1.0F : 1.0F;
if (exp == 0xFU && mant == 0x7U) return std::numeric_limits<float>::quiet_NaN();
if (exp == 0U) return sm * (static_cast<float>(mant) * (1.0F / 512.0F));
const float mantissa = 1.0F + static_cast<float>(mant) * (1.0F / 8.0F);
return sm * std::ldexp(mantissa, static_cast<int>(exp) - 7);
}
inline uint8_t F32ToF8E4M3(float f) {
constexpr float kFp8Max = 448.0F;
if (std::isnan(f)) return 0x7FU;
const uint8_t sign = std::signbit(f) ? 0x80U : 0x00U;
const float a = std::fabs(f);
if (!std::isfinite(a) || a >= kFp8Max) return static_cast<uint8_t>(sign | 0x7EU);
if (a == 0.0F) return sign;
int e2 = 0;
const float frac = std::frexp(a, &e2);
int exp_field = (e2 - 1) + 7;
if (exp_field <= 0) {
const double qd = static_cast<double>(a) * 512.0;
const int qi = static_cast<int>(std::nearbyint(qd));
if (qi <= 0) return sign;
if (qi < 8) return static_cast<uint8_t>(sign | static_cast<uint8_t>(qi));
return static_cast<uint8_t>(sign | (1U << 3));
}
const double sig = static_cast<double>(frac) * 2.0;
int mi = static_cast<int>(std::nearbyint(sig * 8.0));
if (mi == 16) {
mi = 8;
exp_field += 1;
}
const int mant = mi - 8;
if (exp_field > 15 || (exp_field == 15 && mant >= 7)) {
return static_cast<uint8_t>(sign | 0x7EU);
}
return static_cast<uint8_t>(sign | (static_cast<uint8_t>(exp_field) << 3) |
static_cast<uint8_t>(mant));
}
inline uint8_t StoreKvFp8E4M3(float hp, float scale) {
return F32ToF8E4M3(hp / scale);
}
inline float LoadKvFp8E4M3(uint8_t byte, float scale) { return F8E4M3ToF32(byte) * scale; }
}
#endif