llama + CUDA: flash_attn_ext_rows (fix unified kv for multi-seq)

This commit is contained in:
Aman Gupta
2026-09-27 15:32:38 +08:00
parent 85ca3b52c3
commit 7908e9e8ce
27 changed files with 626 additions and 68 deletions
+108
View File
@@ -7994,6 +7994,102 @@ struct test_flash_attn_ext : public test_case {
}
};
// GGML_OP_FLASH_ATTN_EXT with kv_rows: each q slice attends to its own list of K/V cache rows
struct test_flash_attn_ext_kv_rows : public test_case {
const int64_t hs;
const int64_t nh; // K/V heads
const int64_t gqa;
const int64_t n_cells; // rows of the K/V cache
const int64_t kv; // list length per slice
const int64_t nb; // queries per slice
const int64_t n_slices;
const ggml_type type_KV;
std::vector<int32_t> rows;
std::string vars() override {
return VARS_TO_STR8(hs, nh, gqa, n_cells, kv, nb, n_slices, type_KV);
}
double max_nmse_err() override {
return 5e-4;
}
test_flash_attn_ext_kv_rows(int64_t hs = 256, int64_t nh = 2, int64_t gqa = 8, int64_t n_cells = 2048, int64_t kv = 512,
int64_t nb = 8, int64_t n_slices = 4, ggml_type type_KV = GGML_TYPE_F16)
: hs(hs), nh(nh), gqa(gqa), n_cells(n_cells), kv(kv), nb(nb), n_slices(n_slices), type_KV(type_KV) {}
ggml_tensor * build_graph(ggml_context * ctx) override {
ggml_tensor * q = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, hs, nb, nh*gqa, n_slices);
ggml_set_name(q, "q");
// cache layout [hs, nh, n_cells], shared by all slices
ggml_tensor * k0 = ggml_new_tensor_3d(ctx, type_KV, hs, nh, n_cells);
ggml_tensor * v0 = ggml_new_tensor_3d(ctx, type_KV, hs, nh, n_cells);
ggml_set_name(k0, "k0");
ggml_set_name(v0, "v0");
ggml_tensor * k = ggml_permute(ctx, k0, 0, 2, 1, 3);
ggml_tensor * v = ggml_permute(ctx, v0, 0, 2, 1, 3);
ggml_tensor * m = ggml_new_tensor_4d(ctx, GGML_TYPE_F16, kv, nb, 1, n_slices);
ggml_set_name(m, "m");
ggml_tensor * r = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, kv, n_slices);
ggml_set_name(r, "r");
ggml_tensor * out = ggml_flash_attn_ext_rows(ctx, q, k, v, m, r, 1.0f/sqrtf(hs), 0.0f, 0.0f);
ggml_prec_set_acc(out, GGML_PREC_F32);
ggml_set_name(out, "out");
return out;
}
void initialize_tensors(ggml_context * ctx) override {
std::mt19937 gen(0x6B76);
// per slice: distinct random cells, the tail of the list is padding
rows.assign(kv*n_slices, -1);
std::vector<int32_t> order(n_cells);
for (int64_t s = 0; s < n_slices; ++s) {
for (int64_t i = 0; i < n_cells; ++i) {
order[i] = i;
}
std::shuffle(order.begin(), order.end(), gen);
const int64_t n_used = std::min<int64_t>(n_cells, kv - (s*37) % std::max<int64_t>(1, kv/2));
for (int64_t i = 0; i < n_used; ++i) {
rows[s*kv + i] = order[i];
}
}
for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != NULL; t = ggml_get_next_tensor(ctx, t)) {
if (t->view_src) {
continue;
}
if (strcmp(t->name, "r") == 0) {
ggml_backend_tensor_set(t, rows.data(), 0, rows.size()*sizeof(int32_t));
} else if (strcmp(t->name, "m") == 0) {
std::vector<float> f(ggml_nelements(t));
std::uniform_real_distribution<float> dis(-1.0f, 1.0f);
for (int64_t s = 0; s < n_slices; ++s) {
for (int64_t j = 0; j < nb; ++j) {
for (int64_t i = 0; i < kv; ++i) {
const bool pad = rows[s*kv + i] < 0;
const bool drop = (i + 3*j + 5*s) % 11 == 0;
f[(s*nb + j)*kv + i] = pad || drop ? -INFINITY : dis(gen);
}
}
}
std::vector<ggml_fp16_t> h(f.size());
ggml_fp32_to_fp16_row(f.data(), h.data(), f.size());
ggml_backend_tensor_set(t, h.data(), 0, h.size()*sizeof(ggml_fp16_t));
} else {
init_tensor_uniform(t);
}
}
}
};
// GGML_OP_CROSS_ENTROPY_LOSS
struct test_cross_entropy_loss : public test_case {
const ggml_type type;
@@ -11013,6 +11109,18 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
test_cases.emplace_back(new test_flash_attn_ext(128, 128, 1, { 8, 1}, 4096, 4, true, false, 8.0f, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 512));
// flash attention over a list of K/V rows per slice (unified KV cache)
for (int64_t nb : {1, 3, 8, 64, 512}) {
for (int64_t n_slices : {1, 4}) {
test_cases.emplace_back(new test_flash_attn_ext_kv_rows(256, 2, 8, 4096, 1024, nb, n_slices));
test_cases.emplace_back(new test_flash_attn_ext_kv_rows(128, 8, 4, 4096, 512, nb, n_slices));
test_cases.emplace_back(new test_flash_attn_ext_kv_rows(128, 4, 1, 1000, 256, nb, n_slices));
test_cases.emplace_back(new test_flash_attn_ext_kv_rows( 64, 4, 2, 3000, 768, nb, n_slices, GGML_TYPE_Q8_0));
test_cases.emplace_back(new test_flash_attn_ext_kv_rows(512, 2, 8, 4096, 1024, nb, n_slices));
test_cases.emplace_back(new test_flash_attn_ext_kv_rows(256, 8, 2, 4096, 1280, nb, n_slices));
}
}
// sparse mask with large batch size
test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, { 8, 1}, 4096, 64, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 512));
test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, { 8, 1}, 4096, 64, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 2048));