From 574eb984df50f8d6850df20b02ad7c36ad13fda5 Mon Sep 17 00:00:00 2001 From: Oliver Simons Date: Thu, 1 Oct 2026 17:51:48 +0200 Subject: [PATCH] Add supported for strided copies to/from FP8 in CPU --- ggml/src/ggml-cpu/common.h | 14 ++++++++++++++ ggml/src/ggml-cpu/ops.cpp | 10 +++++++++- tests/test-backend-ops.cpp | 10 ++++++++++ 3 files changed, 33 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-cpu/common.h b/ggml/src/ggml-cpu/common.h index abbadc359c..2c62e6976f 100644 --- a/ggml/src/ggml-cpu/common.h +++ b/ggml/src/ggml-cpu/common.h @@ -43,6 +43,14 @@ static inline float f32_to_f32(float x) { return x; } +static inline float f8_e4m3_to_f32(ggml_fp8_e4m3_t x) { + return ggml_f8_e4m3_to_fp32(x.bits); +} + +static inline ggml_fp8_e4m3_t f32_to_f8_e4m3(float x) { + return { ggml_fp32_to_f8_e4m3(x) }; +} + // TODO - merge this into the traits table, after using row-based conversions template struct type_conversion_table; @@ -65,6 +73,12 @@ struct type_conversion_table { static constexpr ggml_bf16_t (*from_f32)(float) = f32_to_bf16; }; +template <> +struct type_conversion_table { + static constexpr float (*to_f32)(ggml_fp8_e4m3_t) = f8_e4m3_to_f32; + static constexpr ggml_fp8_e4m3_t (*from_f32)(float) = f32_to_f8_e4m3; +}; + template <> struct type_conversion_table { static constexpr float (*to_f32)(int32_t) = i32_to_f32; diff --git a/ggml/src/ggml-cpu/ops.cpp b/ggml/src/ggml-cpu/ops.cpp index 69b7e02878..69cb3a5f6f 100644 --- a/ggml/src/ggml-cpu/ops.cpp +++ b/ggml/src/ggml-cpu/ops.cpp @@ -540,6 +540,7 @@ void ggml_compute_forward_dup( /**/ if (dst->type == GGML_TYPE_F16) ggml_compute_forward_dup_flt(params, dst); else if (dst->type == GGML_TYPE_BF16) ggml_compute_forward_dup_flt(params, dst); else if (dst->type == GGML_TYPE_F32) ggml_compute_forward_dup_flt(params, dst); + else if (dst->type == GGML_TYPE_F8_E4M3) ggml_compute_forward_dup_flt(params, dst); else ggml_compute_forward_dup_to_q(params, dst); } break; case GGML_TYPE_BF16: @@ -547,6 +548,7 @@ void ggml_compute_forward_dup( /**/ if (dst->type == GGML_TYPE_F16) ggml_compute_forward_dup_flt(params, dst); else if (dst->type == GGML_TYPE_BF16) ggml_compute_forward_dup_flt(params, dst); else if (dst->type == GGML_TYPE_F32) ggml_compute_forward_dup_flt(params, dst); + else if (dst->type == GGML_TYPE_F8_E4M3) ggml_compute_forward_dup_flt(params, dst); else ggml_compute_forward_dup_to_q(params, dst); } break; case GGML_TYPE_F32: @@ -554,9 +556,15 @@ void ggml_compute_forward_dup( /**/ if (dst->type == GGML_TYPE_F16) ggml_compute_forward_dup_flt(params, dst); else if (dst->type == GGML_TYPE_BF16) ggml_compute_forward_dup_flt(params, dst); else if (dst->type == GGML_TYPE_F32) ggml_compute_forward_dup_flt(params, dst); + else if (dst->type == GGML_TYPE_F8_E4M3) ggml_compute_forward_dup_flt(params, dst); else if (dst->type == GGML_TYPE_I32) ggml_compute_forward_dup_flt(params, dst); else ggml_compute_forward_dup_to_q(params, dst); } break; + case GGML_TYPE_F8_E4M3: + { + if (dst->type == GGML_TYPE_F32) ggml_compute_forward_dup_flt(params, dst); + else GGML_ABORT("not implemented"); + } break; case GGML_TYPE_I32: { if (dst->type == GGML_TYPE_F32) ggml_compute_forward_dup_flt(params, dst); @@ -564,7 +572,7 @@ void ggml_compute_forward_dup( } break; default: { - if ((ggml_is_quantized(src0->type) || src0->type == GGML_TYPE_F8_E4M3) && dst->type == GGML_TYPE_F32) { + if (ggml_is_quantized(src0->type) && dst->type == GGML_TYPE_F32) { ggml_compute_forward_dup_from_q(params, dst); break; } diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 889e28881f..e21f817e4c 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -9804,6 +9804,16 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_cpy(type_src, type_dst, {256, 2, 3, 4}, {-1,-1,-1,-1}, {1, 0, 2, 3})); // cpy not-contiguous } } + for (ggml_type type_src : {GGML_TYPE_F32, GGML_TYPE_F16, GGML_TYPE_BF16}) { + test_cases.emplace_back(new test_cpy(type_src, GGML_TYPE_F8_E4M3, {32, 2, 3, 4}, {-1,-1,-1,-1}, {0, 0, 0, 0}, {0, 0, 0, 0}, false, {64, 2, 3, 4})); + test_cases.emplace_back(new test_cpy(type_src, GGML_TYPE_F8_E4M3, {32, 2, 3, 4}, {-1,-1,-1,-1}, {1, 0, 2, 3}, {1, 0, 2, 3})); + test_cases.emplace_back(new test_cpy(type_src, GGML_TYPE_F8_E4M3, {32, 2, 3, 4}, {16, 4, 3, 4}, {1, 0, 2, 3}, {0, 2, 1, 3})); + } + + test_cases.emplace_back(new test_cpy(GGML_TYPE_F8_E4M3, GGML_TYPE_F32, {32, 2, 3, 4}, {-1,-1,-1,-1}, {0, 0, 0, 0}, {0, 0, 0, 0}, false, {64, 2, 3, 4})); + test_cases.emplace_back(new test_cpy(GGML_TYPE_F8_E4M3, GGML_TYPE_F32, {32, 2, 3, 4}, {-1,-1,-1,-1}, {1, 0, 2, 3}, {1, 0, 2, 3})); + test_cases.emplace_back(new test_cpy(GGML_TYPE_F8_E4M3, GGML_TYPE_F32, {32, 2, 3, 4}, {16, 4, 3, 4}, {1, 0, 2, 3}, {0, 2, 1, 3})); + // quant block count not a multiple of the kernel block size test_cases.emplace_back(new test_cpy(GGML_TYPE_F32, GGML_TYPE_Q4_0, {96, 1, 1, 1})); test_cases.emplace_back(new test_cpy(GGML_TYPE_Q4_0, GGML_TYPE_F32, {96, 1, 1, 1}));