mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-03 03:17:32 -05:00
Add supported for strided copies to/from FP8 in CPU
This commit is contained in:
@@ -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;
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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}));
|
||||
|
||||
Reference in New Issue
Block a user