Perform indexer score calculation in a loop over attention heads

This commit is contained in:
Kawrakow
2026-06-29 17:18:25 +00:00
parent 03391cb5fa
commit d536359431
+41 -19
View File
@@ -463,34 +463,56 @@ ggml_tensor * llm_build_context::build_deepseek2_dsa_indexer(
// ---- scores ----
// indexer_q : {head_size, n_ihead, n_tokens} -> {head_size, n_tokens, n_ihead}
indexer_q = ggml_cont(ctx0, ggml_permute(ctx0, indexer_q, 0, 2, 1, 3));
//indexer_q = ggml_cont(ctx0, ggml_permute(ctx0, indexer_q, 0, 2, 1, 3));
// cached_k : {head_size, n_kv} -> {head_size, n_kv, 1}; broadcasts over q's n_ihead dim.
ggml_tensor * indexer_k_b = ggml_reshape_3d(ctx0, cached_k, head_size, n_kv, 1);
// {n_kv(keys), n_tokens(q), n_ihead} (k's head dim broadcasts over n_ihead)
ggml_tensor * indexer_kq = ggml_mul_mat(ctx0, indexer_k_b, indexer_q);
cb(indexer_kq, "dsa_indexer_kq", il);
ggml_tensor * indexer_score = ggml_view_2d(ctx0, KQ_mask, n_kv, n_tokens, KQ_mask->nb[1], 0);
if (indexer_score->type != GGML_TYPE_F32) {
indexer_score = ggml_cast(ctx0, indexer_score, GGML_TYPE_F32);
}
for (int head = 0; head < n_ihead; ++head) {
// [head_size, n_tokens]
auto q = ggml_view_2d(ctx0, indexer_q, indexer_q->ne[0], indexer_q->ne[2], indexer_q->nb[2], indexer_q->nb[1]*head);
// [n_kv, n_tokens]
auto kq = ggml_mul_mat(ctx0, indexer_k_b, q);
// [n_kv, n_tokens]
kq = ggml_relu(ctx0, kq);
// [1, n_tokens]
auto w = ggml_cont(ctx0, ggml_view_2d(ctx0, indexer_weights, 1, indexer_weights->ne[1], indexer_weights->nb[1], indexer_weights->nb[0]*head));
// [n_kv, n_tokens]
auto score = ggml_mul(ctx0, kq, w);
indexer_score = ggml_add(ctx0, indexer_score, score);
}
// -> {n_ihead, n_tokens(q), n_kv(keys)} for per-head weighting
indexer_kq = ggml_cont(ctx0, ggml_permute(ctx0, indexer_kq, 2, 1, 0, 3));
////// {n_kv(keys), n_tokens(q), n_ihead} (k's head dim broadcasts over n_ihead)
////ggml_tensor * indexer_kq = ggml_mul_mat(ctx0, indexer_k_b, indexer_q);
////cb(indexer_kq, "dsa_indexer_kq", il);
ggml_tensor * indexer_score = ggml_relu(ctx0, indexer_kq);
//// {n_kv(keys), n_tokens(q), n_ihead} (k's head dim broadcasts over n_ihead)
//ggml_tensor * indexer_kq = ggml_mul_mat(ctx0, indexer_k_b, indexer_q);
//cb(indexer_kq, "dsa_indexer_kq", il);
// weights {n_ihead, n_tokens} -> {n_ihead, n_tokens, 1} broadcast over keys
indexer_weights = ggml_reshape_3d(ctx0, indexer_weights, n_ihead, n_tokens, 1);
indexer_score = ggml_mul(ctx0, indexer_score, indexer_weights);
//// -> {n_ihead, n_tokens(q), n_kv(keys)} for per-head weighting
//indexer_kq = ggml_cont(ctx0, ggml_permute(ctx0, indexer_kq, 2, 1, 0, 3));
// sum over heads -> {1, n_tokens(q), n_kv(keys)}
indexer_score = ggml_sum_rows(ctx0, indexer_score);
//ggml_tensor * indexer_score = ggml_relu(ctx0, indexer_kq);
// -> {n_kv(keys), n_tokens(q), 1}
indexer_score = ggml_cont(ctx0, ggml_permute(ctx0, indexer_score, 2, 1, 0, 3));
cb(indexer_score, "dsa_indexer_score", il);
//// weights {n_ihead, n_tokens} -> {n_ihead, n_tokens, 1} broadcast over keys
//indexer_weights = ggml_reshape_3d(ctx0, indexer_weights, n_ihead, n_tokens, 1);
//indexer_score = ggml_mul(ctx0, indexer_score, indexer_weights);
// add base causal mask over the n_kv keys: first n_tokens query columns of KQ_mask {n_kv, n_tokens_pad}.
ggml_tensor * causal = ggml_view_2d(ctx0, KQ_mask, n_kv, n_tokens, KQ_mask->nb[1], 0);
indexer_score = ggml_add(ctx0, indexer_score, causal);
cb(indexer_score, "dsa_indexer_score_masked", il);
//// sum over heads -> {1, n_tokens(q), n_kv(keys)}
//indexer_score = ggml_sum_rows(ctx0, indexer_score);
//// -> {n_kv(keys), n_tokens(q), 1}
//indexer_score = ggml_cont(ctx0, ggml_permute(ctx0, indexer_score, 2, 1, 0, 3));
//cb(indexer_score, "dsa_indexer_score", il);
//// add base causal mask over the n_kv keys: first n_tokens query columns of KQ_mask {n_kv, n_tokens_pad}.
//ggml_tensor * causal = ggml_view_2d(ctx0, KQ_mask, n_kv, n_tokens, KQ_mask->nb[1], 0);
//indexer_score = ggml_add(ctx0, indexer_score, causal);
//cb(indexer_score, "dsa_indexer_score_masked", il);
// Attention-sink force-inclusion: add a finite positive boost to each query's OWN SEQUENCE's
// first n_sink present tokens so the sink token(s) always survive the top-k selection. Masking