mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-02 19:07:25 -05:00
llama + CUDA: flash_attn_ext_rows (fix unified kv for multi-seq)
This commit is contained in:
@@ -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));
|
||||
|
||||
Reference in New Issue
Block a user