Add supported for strided copies to/from FP8 in CPU

This commit is contained in:
Oliver Simons
2026-10-01 17:59:18 +02:00
parent aa742bb620
commit 574eb984df
3 changed files with 33 additions and 1 deletions
+14
View File
@@ -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 <class T>
struct type_conversion_table;
@@ -65,6 +73,12 @@ struct type_conversion_table<ggml_bf16_t> {
static constexpr ggml_bf16_t (*from_f32)(float) = f32_to_bf16;
};
template <>
struct type_conversion_table<ggml_fp8_e4m3_t> {
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<int32_t> {
static constexpr float (*to_f32)(int32_t) = i32_to_f32;
+9 -1
View File
@@ -540,6 +540,7 @@ void ggml_compute_forward_dup(
/**/ if (dst->type == GGML_TYPE_F16) ggml_compute_forward_dup_flt<ggml_fp16_t, ggml_fp16_t>(params, dst);
else if (dst->type == GGML_TYPE_BF16) ggml_compute_forward_dup_flt<ggml_fp16_t, ggml_bf16_t>(params, dst);
else if (dst->type == GGML_TYPE_F32) ggml_compute_forward_dup_flt<ggml_fp16_t, float >(params, dst);
else if (dst->type == GGML_TYPE_F8_E4M3) ggml_compute_forward_dup_flt<ggml_fp16_t, ggml_fp8_e4m3_t>(params, dst);
else ggml_compute_forward_dup_to_q<ggml_fp16_t>(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<ggml_bf16_t, ggml_fp16_t>(params, dst);
else if (dst->type == GGML_TYPE_BF16) ggml_compute_forward_dup_flt<ggml_bf16_t, ggml_bf16_t>(params, dst);
else if (dst->type == GGML_TYPE_F32) ggml_compute_forward_dup_flt<ggml_bf16_t, float >(params, dst);
else if (dst->type == GGML_TYPE_F8_E4M3) ggml_compute_forward_dup_flt<ggml_bf16_t, ggml_fp8_e4m3_t>(params, dst);
else ggml_compute_forward_dup_to_q<ggml_bf16_t>(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<float, ggml_fp16_t>(params, dst);
else if (dst->type == GGML_TYPE_BF16) ggml_compute_forward_dup_flt<float, ggml_bf16_t>(params, dst);
else if (dst->type == GGML_TYPE_F32) ggml_compute_forward_dup_flt<float, float >(params, dst);
else if (dst->type == GGML_TYPE_F8_E4M3) ggml_compute_forward_dup_flt<float, ggml_fp8_e4m3_t>(params, dst);
else if (dst->type == GGML_TYPE_I32) ggml_compute_forward_dup_flt<float, int32_t >(params, dst);
else ggml_compute_forward_dup_to_q<float>(params, dst);
} break;
case GGML_TYPE_F8_E4M3:
{
if (dst->type == GGML_TYPE_F32) ggml_compute_forward_dup_flt<ggml_fp8_e4m3_t, float>(params, dst);
else GGML_ABORT("not implemented");
} break;
case GGML_TYPE_I32:
{
if (dst->type == GGML_TYPE_F32) ggml_compute_forward_dup_flt<int32_t, float>(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;
}
+10
View File
@@ -9804,6 +9804,16 @@ static std::vector<std::unique_ptr<test_case>> 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}));