DSV4: add shared top-k selection and improve mask handling

This commit is contained in:
samuel
2026-07-12 19:03:57 -03:00
parent 41bc2e7afd
commit 83b3e145b5
3 changed files with 69 additions and 9 deletions
+10
View File
@@ -154,11 +154,21 @@ void ggml_cuda_op_concat(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
const ggml_tensor * src1 = dst->src[1];
const int32_t dim = ((int32_t *) dst->op_params)[0];
GGML_ASSERT(dim >= 0 && dim < GGML_MAX_DIMS);
GGML_ASSERT(src0->type == src1->type);
GGML_ASSERT(dst->type == src0->type);
GGML_ASSERT(!ggml_is_quantized(src0->type));
GGML_ASSERT(ggml_blck_size(src0->type) == 1);
for (int d = 0; d < GGML_MAX_DIMS; ++d) {
if (d == dim) {
GGML_ASSERT(dst->ne[d] == src0->ne[d] + src1->ne[d]);
} else {
GGML_ASSERT(src0->ne[d] == src1->ne[d]);
GGML_ASSERT(dst->ne[d] == src0->ne[d]);
}
}
switch (ggml_type_size(src0->type)) {
case 1: concat_cuda<uint8_t>(src0, src1, dst, dim, ctx.stream()); break;
case 2: concat_cuda<uint16_t>(src0, src1, dst, dim, ctx.stream()); break;
+8 -7
View File
@@ -10197,12 +10197,18 @@ static struct ggml_tensor * ggml_fill_impl(
GGML_ASSERT(a->type == GGML_TYPE_F32 || a->type == GGML_TYPE_F16);
GGML_ASSERT(ggml_is_contiguous(a));
bool is_node = false;
if (!inplace && a->grad) {
GGML_ABORT("fatal error"); // TODO: implement backward
is_node = true;
}
struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
ggml_set_op_params_f32(result, 0, c);
result->op = GGML_OP_FILL;
result->grad = NULL;
result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
result->src[0] = a;
return result;
@@ -22209,9 +22215,6 @@ static void ggml_compute_forward_flash_attn_ext_f16(
kq_vec_dot(Dk, &s, 0, k_data, 0, Q_q, 0, 1);
s = softcap == 0.0f ? s*scale + mv : softcap*tanhf(s*scale) + mv; // scale KQ value and apply mask
if (!isfinite(s)) {
continue;
}
const float Mold = M;
@@ -22280,9 +22283,7 @@ static void ggml_compute_forward_flash_attn_ext_f16(
}
// V /= S
// A fully masked query row has no active keys. Match mainline by
// returning a zero attention vector instead of propagating NaNs.
const float S_inv = isfinite(S) && S > 0.0f ? 1.0f/S : 0.0f;
const float S_inv = S == 0.0f ? 0.0f : 1.0f/S;
ggml_vec_scale_f32(Dv, VKQ32, S_inv);
// dst indices
+51 -2
View File
@@ -825,6 +825,45 @@ static ggml_tensor * build_top_k_mask(
return ggml_add(ctx0, kq_mask_top_k, kq_mask);
}
static ggml_tensor * dsv4_build_lid_top_k_shared(
ggml_context * ctx0,
ggml_tensor * indexer_k,
ggml_tensor * indexer_q,
ggml_tensor * indexer_weights,
ggml_tensor * indexer_mask,
int n_top_k) {
const int64_t n_stream = indexer_k->ne[3];
const int64_t n_tokens = indexer_q->ne[1];
if (n_stream <= 0 || indexer_k->ne[2] != 1 || indexer_q->ne[3] != n_stream ||
indexer_weights->ne[3] != n_stream || indexer_mask->ne[1] < n_stream*n_tokens) {
return nullptr;
}
ggml_tensor * selected = nullptr;
for (int64_t s = 0; s < n_stream; ++s) {
ggml_tensor * k = ggml_view_2d(ctx0, indexer_k,
indexer_k->ne[0], indexer_k->ne[1], indexer_k->nb[1], s*indexer_k->nb[3]);
ggml_tensor * q = ggml_view_3d(ctx0, indexer_q,
indexer_q->ne[0], indexer_q->ne[1], indexer_q->ne[2],
indexer_q->nb[1], indexer_q->nb[2], s*indexer_q->nb[3]);
q = ggml_permute(ctx0, q, 0, 2, 1, 3);
ggml_tensor * w = ggml_view_2d(ctx0, indexer_weights,
indexer_weights->ne[0], indexer_weights->ne[1], indexer_weights->nb[1],
s*indexer_weights->nb[3]);
ggml_tensor * mask = ggml_view_2d(ctx0, indexer_mask,
indexer_mask->ne[0], n_tokens, indexer_mask->nb[1],
s*n_tokens*indexer_mask->nb[1]);
ggml_tensor * cur = ggml_indexer_topk(ctx0, k, q, w, mask,
GGML_UNARY_OP_RELU, n_top_k);
selected = selected == nullptr ? cur : ggml_concat(ctx0, selected, cur, 1);
}
return selected == nullptr ? nullptr : ggml_cont(ctx0, selected);
}
static ggml_tensor * dsv4_build_lid_top_k(
ggml_context * ctx0,
llm_build_context & llm,
@@ -890,6 +929,17 @@ static ggml_tensor * dsv4_build_lid_top_k(
indexer_k = ggml_permute(ctx0, indexer_k, 0, 2, 1, 3);
llm.cb(indexer_k, "lid_k_stream", il);
ggml_tensor * lid_mask = dsv4_build_raw_mask_view(ctx0,
llm.lctx.dsv4.inputs.lid.kq_mask, n_lid, n_tokens);
const uint32_t n_top_k = (uint32_t) std::min<int64_t>(n_lid, hparams.indexer_top_k);
if (llm.cparams.fused_idx_topk && n_lid > n_top_k) {
if (ggml_tensor * selected = dsv4_build_lid_top_k_shared(ctx0,
indexer_k, indexer_q, indexer_weights, lid_mask, (int) n_top_k)) {
llm.cb(selected, "lid_top_k", il);
return selected;
}
}
ggml_tensor * indexer_kq = ggml_mul_mat(ctx0, indexer_k, indexer_q);
llm.cb(indexer_kq, "lid_kq", il);
@@ -903,10 +953,9 @@ static ggml_tensor * dsv4_build_lid_top_k(
indexer_score = ggml_view_2d(ctx0, indexer_score, n_lid, n_tokens, indexer_score->nb[2], 0);
llm.cb(indexer_score, "lid_score", il);
indexer_score = ggml_add(ctx0, indexer_score, dsv4_build_raw_mask_view(ctx0, llm.lctx.dsv4.inputs.lid.kq_mask, n_lid, n_tokens));
indexer_score = ggml_add(ctx0, indexer_score, lid_mask);
llm.cb(indexer_score, "lid_score_masked", il);
const uint32_t n_top_k = (uint32_t) std::min<int64_t>(indexer_score->ne[0], hparams.indexer_top_k);
ggml_tensor * top_k = ggml_cont(ctx0, ggml_top_k(ctx0, indexer_score, n_top_k));
llm.cb(top_k, "lid_top_k", il);