Add tests for scaled, FP8 KV-cache to test-backend-ops

This commit is contained in:
Oliver Simons
2026-10-09 21:07:43 +02:00
parent 29995854dc
commit 287791c289
+77 -8
View File
@@ -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 kv_view; // create K/V as views of a larger buffer (like a KV cache)
const bool v_is_view_of_k; const bool v_is_view_of_k;
const int64_t n_kv_max; const int64_t n_kv_max;
const float scale_K;
const float scale_V;
const bool flash;
std::string vars() override { 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 { 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, 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, 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}, 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), : 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 { ggml_tensor * build_graph(ggml_context * ctx) override {
const int64_t hsk_padded = GGML_PAD(hsk, ggml_blck_size(type_K)); 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_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"); 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 = nullptr;
ggml_tensor * v_scale = nullptr;
if (v_is_view_of_k) { 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 // the V cache is a sub-view of the K cache. this is used by some MLA-based models
// for more info: // 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); v = ggml_view_4d(ctx, k, hsv_padded, kv, nh, nr23[1], k->nb[1], k->nb[2], k->nb[3], 0);
} else { } 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"); 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; ggml_tensor * m = nullptr;
if (mask) { if (mask) {
@@ -8435,10 +8472,21 @@ struct test_flash_attn_ext : public test_case {
ggml_set_name(s, "s"); 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_tensor * out;
ggml_flash_attn_ext_add_sinks(out, s); if (flash) {
ggml_flash_attn_ext_set_n_kv_max(out, n_kv_max); out = ggml_flash_attn_ext(ctx, q, k, v, m, 1.0f/sqrtf(hsk), max_bias, logit_softcap);
ggml_prec_set_acc(out, prec); 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"); ggml_set_name(out, "out");
return out; return out;
@@ -8455,6 +8503,16 @@ struct test_flash_attn_ext : public test_case {
} else { } else {
init_tensor_kq_mask(t); 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 { } else {
init_tensor_uniform(t); 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_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_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)); 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 // 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)); 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));