mirror of
https://github.com/ikawrakow/ik_llama.cpp.git
synced 2026-07-21 10:15:35 +00:00
DSV4: add shared top-k selection and improve mask handling
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user