mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-11 15:30:39 +02:00
Add tests for scaled, FP8 KV-cache to test-backend-ops
This commit is contained in:
@@ -8356,9 +8356,23 @@ struct test_flash_attn_ext : public test_case {
|
||||
const bool kv_view; // create K/V as views of a larger buffer (like a KV cache)
|
||||
const bool v_is_view_of_k;
|
||||
const int64_t n_kv_max;
|
||||
const float scale_K;
|
||||
const float scale_V;
|
||||
const bool flash;
|
||||
|
||||
std::string vars() override {
|
||||
return VARS_TO_STR17(hsk, hsv, nh, nr23, kv, nb, mask, sinks, max_bias, logit_softcap, prec, type_K, type_V, permute, kv_view, v_is_view_of_k, n_kv_max);
|
||||
std::string result = VARS_TO_STR17(hsk, hsv, nh, nr23, kv, nb, mask, sinks, max_bias, logit_softcap, prec, type_K, type_V, permute, kv_view, v_is_view_of_k, n_kv_max);
|
||||
if (scale_K != 1.0f || scale_V != 1.0f) {
|
||||
result += "," + VARS_TO_STR3(scale_K, scale_V, flash);
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
std::string op_desc(ggml_tensor * t) override {
|
||||
if (scale_K != 1.0f || scale_V != 1.0f) {
|
||||
return flash ? "FLASH_ATTN_EXT_SCALED" : "MUL_MAT_SCALED_KV";
|
||||
}
|
||||
return test_case::op_desc(t);
|
||||
}
|
||||
|
||||
double max_nmse_err() override {
|
||||
@@ -8375,9 +8389,10 @@ struct test_flash_attn_ext : public test_case {
|
||||
test_flash_attn_ext(int64_t hsk = 128, int64_t hsv = 128, int64_t nh = 32, std::array<int64_t, 2> nr23 = {1, 1}, int64_t kv = 96, int64_t nb = 8,
|
||||
bool mask = true, bool sinks = false, float max_bias = 0.0f, float logit_softcap = 0.0f, ggml_prec prec = GGML_PREC_F32,
|
||||
ggml_type type_K = GGML_TYPE_F16, ggml_type type_V = GGML_TYPE_F16, std::array<int32_t, 4> permute = {0, 1, 2, 3},
|
||||
bool kv_view = true, bool v_is_view_of_k = false, int64_t n_kv_max = 0)
|
||||
bool kv_view = true, bool v_is_view_of_k = false, int64_t n_kv_max = 0, float scale_K = 1.0f, float scale_V = 1.0f, bool flash = true)
|
||||
: hsk(hsk), hsv(hsv), nh(nh), nr23(nr23), kv(kv), nb(nb), mask(mask), sinks(sinks), max_bias(max_bias), logit_softcap(logit_softcap), prec(prec),
|
||||
type_K(type_K), type_V(type_V), permute(permute), kv_view(kv_view), v_is_view_of_k(v_is_view_of_k), n_kv_max(n_kv_max) {}
|
||||
type_K(type_K), type_V(type_V), permute(permute), kv_view(kv_view), v_is_view_of_k(v_is_view_of_k), n_kv_max(n_kv_max),
|
||||
scale_K(scale_K), scale_V(scale_V), flash(flash) {}
|
||||
|
||||
ggml_tensor * build_graph(ggml_context * ctx) override {
|
||||
const int64_t hsk_padded = GGML_PAD(hsk, ggml_blck_size(type_K));
|
||||
@@ -8407,8 +8422,20 @@ struct test_flash_attn_ext : public test_case {
|
||||
|
||||
ggml_tensor * k = create_permuted(type_K, hsk_padded, kv, nh, nr23[1], kv_view); // the K tensor is usually a view of the K cache
|
||||
ggml_set_name(k, "k");
|
||||
if (scale_K != 1.0f) {
|
||||
GGML_ASSERT(type_K == GGML_TYPE_F8_E4M3);
|
||||
ggml_tensor * k_cur = create_permuted(GGML_TYPE_F32, hsk_padded, kv, nh, nr23[1], false);
|
||||
ggml_set_name(k_cur, "k_cur");
|
||||
ggml_tensor * k_scale = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 1);
|
||||
ggml_set_name(k_scale, "k_scale");
|
||||
ggml_tensor * k_idxs = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, kv);
|
||||
ggml_set_name(k_idxs, "k_idxs");
|
||||
k = ggml_set_rows(ctx, k, ggml_div(ctx, k_cur, k_scale), k_idxs);
|
||||
q = ggml_mul(ctx, q, k_scale);
|
||||
}
|
||||
|
||||
ggml_tensor * v = nullptr;
|
||||
ggml_tensor * v_scale = nullptr;
|
||||
if (v_is_view_of_k) {
|
||||
// the V cache is a sub-view of the K cache. this is used by some MLA-based models
|
||||
// for more info:
|
||||
@@ -8419,9 +8446,19 @@ struct test_flash_attn_ext : public test_case {
|
||||
|
||||
v = ggml_view_4d(ctx, k, hsv_padded, kv, nh, nr23[1], k->nb[1], k->nb[2], k->nb[3], 0);
|
||||
} else {
|
||||
v = create_permuted(type_V, hsv_padded, kv, nh, nr23[1], kv_view); // the V tensor is usually a view of the V cache
|
||||
v = create_permuted(type_V, flash ? hsv_padded : kv, flash ? kv : hsv_padded, nh, nr23[1], kv_view); // the V tensor is usually a view of the V cache
|
||||
}
|
||||
ggml_set_name(v, "v");
|
||||
if (scale_V != 1.0f) {
|
||||
GGML_ASSERT(type_V == GGML_TYPE_F8_E4M3 && !v_is_view_of_k);
|
||||
ggml_tensor * v_cur = create_permuted(GGML_TYPE_F32, flash ? hsv_padded : kv, flash ? kv : hsv_padded, nh, nr23[1], false);
|
||||
ggml_set_name(v_cur, "v_cur");
|
||||
v_scale = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 1);
|
||||
ggml_set_name(v_scale, "v_scale");
|
||||
ggml_tensor * v_idxs = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, v->ne[1]);
|
||||
ggml_set_name(v_idxs, "v_idxs");
|
||||
v = ggml_set_rows(ctx, v, ggml_div(ctx, v_cur, v_scale), v_idxs);
|
||||
}
|
||||
|
||||
ggml_tensor * m = nullptr;
|
||||
if (mask) {
|
||||
@@ -8435,10 +8472,21 @@ struct test_flash_attn_ext : public test_case {
|
||||
ggml_set_name(s, "s");
|
||||
}
|
||||
|
||||
ggml_tensor * out = ggml_flash_attn_ext(ctx, q, k, v, m, 1.0f/sqrtf(hsk), max_bias, logit_softcap);
|
||||
ggml_flash_attn_ext_add_sinks(out, s);
|
||||
ggml_flash_attn_ext_set_n_kv_max(out, n_kv_max);
|
||||
ggml_prec_set_acc(out, prec);
|
||||
ggml_tensor * out;
|
||||
if (flash) {
|
||||
out = ggml_flash_attn_ext(ctx, q, k, v, m, 1.0f/sqrtf(hsk), max_bias, logit_softcap);
|
||||
ggml_flash_attn_ext_add_sinks(out, s);
|
||||
ggml_flash_attn_ext_set_n_kv_max(out, n_kv_max);
|
||||
ggml_prec_set_acc(out, prec);
|
||||
} else {
|
||||
GGML_ASSERT(!sinks && max_bias == 0.0f && logit_softcap == 0.0f);
|
||||
ggml_tensor * kq = ggml_mul_mat(ctx, k, q);
|
||||
ggml_prec_set_acc(kq, prec);
|
||||
out = ggml_mul_mat(ctx, v, ggml_soft_max_ext(ctx, kq, m, 1.0f/sqrtf(hsk), 0.0f));
|
||||
}
|
||||
if (v_scale) {
|
||||
out = ggml_mul(ctx, out, v_scale);
|
||||
}
|
||||
ggml_set_name(out, "out");
|
||||
|
||||
return out;
|
||||
@@ -8455,6 +8503,16 @@ struct test_flash_attn_ext : public test_case {
|
||||
} else {
|
||||
init_tensor_kq_mask(t);
|
||||
}
|
||||
} else if (strcmp(t->name, "k_scale") == 0) {
|
||||
ggml_backend_tensor_set(t, &scale_K, 0, sizeof(scale_K));
|
||||
} else if (strcmp(t->name, "v_scale") == 0) {
|
||||
ggml_backend_tensor_set(t, &scale_V, 0, sizeof(scale_V));
|
||||
} else if (strcmp(t->name, "k_idxs") == 0 || strcmp(t->name, "v_idxs") == 0) {
|
||||
std::vector<int32_t> idxs(t->ne[0]);
|
||||
for (int32_t i = 0; i < t->ne[0]; ++i) {
|
||||
idxs[i] = i;
|
||||
}
|
||||
ggml_backend_tensor_set(t, idxs.data(), 0, idxs.size()*sizeof(int32_t));
|
||||
} else {
|
||||
init_tensor_uniform(t);
|
||||
}
|
||||
@@ -11660,6 +11718,17 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
|
||||
test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 128, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F8_E4M3, GGML_TYPE_F8_E4M3));
|
||||
test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 128, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F8_E4M3, GGML_TYPE_F16));
|
||||
test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 128, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F8_E4M3));
|
||||
for (bool flash : {false, true}) {
|
||||
for (int nb : {1, 8}) {
|
||||
for (int scaled : {1, 2, 3}) {
|
||||
const bool scale_k = scaled & 1;
|
||||
const bool scale_v = scaled & 2;
|
||||
test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 128, nb, true, false, 0, 0, GGML_PREC_F32,
|
||||
scale_k ? GGML_TYPE_F8_E4M3 : GGML_TYPE_F16, scale_v ? GGML_TYPE_F8_E4M3 : GGML_TYPE_F16,
|
||||
{0, 1, 2, 3}, true, false, 0, scale_k ? 0.03125f : 1.0f, scale_v ? 0.125f : 1.0f, flash));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// q8_0 KV cases: decode and prompt batches, KV pad, permuted KV, feature flags, and long context
|
||||
test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {16, 1}, 113, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0));
|
||||
|
||||
Reference in New Issue
Block a user