From 287791c28901776436b303dc928ae04b456c9d5a Mon Sep 17 00:00:00 2001 From: Oliver Simons Date: Wed, 30 Sep 2026 11:31:38 +0200 Subject: [PATCH] Add tests for scaled, FP8 KV-cache to test-backend-ops --- tests/test-backend-ops.cpp | 85 ++++++++++++++++++++++++++++++++++---- 1 file changed, 77 insertions(+), 8 deletions(-) diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 5aa4653ec0..c7b74b19e9 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -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 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 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 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> 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));