mirror of
https://github.com/ikawrakow/ik_llama.cpp.git
synced 2026-07-21 10:15:35 +00:00
WIP
This commit is contained in:
@@ -131,7 +131,8 @@ static ggml_tensor * dsv4_build_raw_mask_view(
|
||||
ggml_tensor * raw_k_read_idxs,
|
||||
int64_t n_kv,
|
||||
int64_t n_tokens,
|
||||
int64_t n_stream) {
|
||||
int64_t n_stream,
|
||||
const llm_build_cb & cb) {
|
||||
const int64_t n_tokens_stream = n_stream > 0 ? n_tokens/n_stream : n_tokens;
|
||||
const int64_t n_rows_stream = GGML_PAD(n_kv, 256);
|
||||
|
||||
@@ -156,6 +157,7 @@ static ggml_tensor * dsv4_build_raw_mask_view(
|
||||
ggml_tensor * stream = ggml_reshape_4d(ctx, ggml_cont(ctx, ggml_transpose(ctx, rows)),
|
||||
n_kv, n_tokens_stream, 1, 1);
|
||||
result = result == nullptr ? stream : ggml_concat(ctx, result, stream, 3);
|
||||
cb(result, "raw_mask_view", s);
|
||||
}
|
||||
return result;
|
||||
}
|
||||
@@ -834,7 +836,7 @@ static ggml_tensor * dsv4_build_lid_top_k_shared(
|
||||
ggml_tensor * indexer_q,
|
||||
ggml_tensor * indexer_weights,
|
||||
ggml_tensor * indexer_mask,
|
||||
int n_top_k) {
|
||||
int n_top_k, const llm_build_cb & cb) {
|
||||
const int64_t n_stream = indexer_k->ne[3];
|
||||
const int64_t n_tokens = indexer_q->ne[1];
|
||||
|
||||
@@ -862,7 +864,13 @@ static ggml_tensor * dsv4_build_lid_top_k_shared(
|
||||
|
||||
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);
|
||||
if (selected) {
|
||||
selected = ggml_concat(ctx0, selected, cur, 1);
|
||||
cb(selected, "top_k", s);
|
||||
} else {
|
||||
selected = cur;
|
||||
}
|
||||
//selected = selected == nullptr ? cur : ggml_concat(ctx0, selected, cur, 1);
|
||||
}
|
||||
|
||||
return selected == nullptr ? nullptr : ggml_cont(ctx0, selected);
|
||||
@@ -874,7 +882,7 @@ static ggml_tensor * dsv4_build_lid_top_k(
|
||||
ggml_tensor * qr,
|
||||
ggml_tensor * cur,
|
||||
ggml_tensor * inp_pos,
|
||||
int il) {
|
||||
int il, ggml_cgraph * gf, const llm_build_cb & cb) {
|
||||
const auto & hparams = llm.hparams;
|
||||
const auto & layer = llm.model.layers[il];
|
||||
const int64_t n_embd_indexer_head = hparams.indexer_head_size;
|
||||
@@ -907,6 +915,7 @@ static ggml_tensor * dsv4_build_lid_top_k(
|
||||
hparams.dsv4_compress_rope_base, llm.freq_scale,
|
||||
llm.ext_factor, dsv4_rope_attn_factor(llm.freq_scale, llm.ext_factor), llm.beta_fast, llm.beta_slow);
|
||||
indexer_q = ggml_concat(ctx0, indexer_q_nope, indexer_q_pe, 0);
|
||||
llm.cb(indexer_q, "indexer_q", il);
|
||||
GGML_ASSERT(indexer_q->ne[0] % hadamard_block == 0);
|
||||
indexer_q = ggml_hadamard(ctx0, indexer_q, hadamard_block);
|
||||
llm.cb(indexer_q, "lid_q_hadamard", il);
|
||||
@@ -937,13 +946,16 @@ static ggml_tensor * dsv4_build_lid_top_k(
|
||||
|
||||
GGML_ASSERT(llm.lctx.dsv4.inputs.csa.kq_mask != nullptr);
|
||||
ggml_tensor * lid_mask = dsv4_build_raw_mask_view(ctx0,
|
||||
llm.lctx.dsv4.inputs.csa.kq_mask, nullptr, n_lid, n_tokens, n_stream);
|
||||
llm.lctx.dsv4.inputs.csa.kq_mask, nullptr, n_lid, n_tokens, n_stream, cb);
|
||||
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;
|
||||
indexer_k, indexer_q, indexer_weights, lid_mask, (int) n_top_k, cb)) {
|
||||
if (selected) {
|
||||
ggml_build_forward_expand(gf, selected);
|
||||
llm.cb(selected, "lid_top_k", il);
|
||||
return selected;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1239,7 +1251,7 @@ ggml_cgraph * llm_build_context::build_deepseek4() {
|
||||
raw_k = dsv4_pad_raw_k_to(ctx0, raw_k, raw_attn_n_kv);
|
||||
}
|
||||
ggml_tensor * raw_mask = dsv4_build_raw_mask_view(ctx0, KQ_mask,
|
||||
lctx.dsv4.inputs.raw_k_read_idxs, raw_kq_n_kv, n_tokens, raw_k->ne[3]);
|
||||
lctx.dsv4.inputs.raw_k_read_idxs, raw_kq_n_kv, n_tokens, raw_k->ne[3], cb);
|
||||
raw_mask = dsv4_pad_mask_tokens(ctx0, raw_mask, n_tokens);
|
||||
raw_mask = dsv4_pad_raw_mask_to(ctx0, raw_mask, raw_attn_n_kv, n_tokens);
|
||||
cb(raw_mask, "dsv4_raw_mask_padded", il);
|
||||
@@ -1255,10 +1267,10 @@ ggml_cgraph * llm_build_context::build_deepseek4() {
|
||||
lctx.dsv4.csa_ctx,
|
||||
n_embd_head,
|
||||
lctx.dsv4.cache.csa_k[(size_t) il]->ne[1]/std::max<uint32_t>(1, lctx.dsv4.cache.n_stream));
|
||||
ggml_tensor * top_k = dsv4_build_lid_top_k(ctx0, *this, qr, cur, inp_pos, il);
|
||||
ggml_tensor * top_k = dsv4_build_lid_top_k(ctx0, *this, qr, cur, inp_pos, il, gf, cb);
|
||||
ggml_tensor * csa_mask = build_top_k_mask(ctx0,
|
||||
dsv4_build_raw_mask_view(ctx0, lctx.dsv4.inputs.csa.kq_mask, nullptr,
|
||||
lctx.dsv4.csa_plan.n_kv, n_tokens, csa_k->ne[3]),
|
||||
lctx.dsv4.csa_plan.n_kv, n_tokens, csa_k->ne[3], cb),
|
||||
top_k);
|
||||
const bool use_fattn = cparams.flash_attn;
|
||||
if (use_fattn) {
|
||||
@@ -1267,7 +1279,7 @@ ggml_cgraph * llm_build_context::build_deepseek4() {
|
||||
raw_k = dsv4_repeat_streams(ctx0, raw_k, csa_k->ne[3]);
|
||||
if (!use_fattn) {
|
||||
raw_mask = dsv4_build_raw_mask_view(ctx0, KQ_mask,
|
||||
lctx.dsv4.inputs.raw_k_read_idxs, raw_kq_n_kv, n_tokens, csa_k->ne[3]);
|
||||
lctx.dsv4.inputs.raw_k_read_idxs, raw_kq_n_kv, n_tokens, csa_k->ne[3], cb);
|
||||
raw_mask = dsv4_pad_raw_mask_to(ctx0, raw_mask, raw_attn_n_kv, n_tokens);
|
||||
}
|
||||
if (use_fattn && csa_mask->type != GGML_TYPE_F16) {
|
||||
@@ -1299,7 +1311,7 @@ ggml_cgraph * llm_build_context::build_deepseek4() {
|
||||
lctx.dsv4.cache.hca_k[(size_t) il]->ne[1]/std::max<uint32_t>(1, lctx.dsv4.cache.n_stream));
|
||||
const bool use_fattn = cparams.flash_attn;
|
||||
ggml_tensor * hca_mask = dsv4_build_raw_mask_view(ctx0, lctx.dsv4.inputs.hca.kq_mask, nullptr,
|
||||
lctx.dsv4.hca_plan.n_kv, n_tokens, hca_k->ne[3]);
|
||||
lctx.dsv4.hca_plan.n_kv, n_tokens, hca_k->ne[3], cb);
|
||||
hca_mask = dsv4_pad_mask_tokens(ctx0, hca_mask, n_tokens);
|
||||
if (use_fattn && hca_mask->type != GGML_TYPE_F16) {
|
||||
hca_mask = ggml_cast(ctx0, hca_mask, GGML_TYPE_F16);
|
||||
@@ -1339,6 +1351,7 @@ ggml_cgraph * llm_build_context::build_deepseek4() {
|
||||
freq_base_l, freq_scale_l, ext_factor_l, attn_factor_l, beta_fast_l, beta_slow_l);
|
||||
cb(attn_pe, "attn_derope", il);
|
||||
attn = ggml_concat(ctx0, attn_nope, attn_pe, 0);
|
||||
cb(attn, "attn", il);
|
||||
|
||||
const int64_t o_group_dim = model.layers[il].wo_a->ne[0];
|
||||
const int64_t n_groups = (n_head * n_embd_head) / o_group_dim;
|
||||
|
||||
@@ -1046,10 +1046,11 @@ ggml_cgraph * llm_build_context::build_openpangu() {
|
||||
h_pre = ggml_add(ctx0, ggml_mul(ctx0, ggml_cont(ctx0, h_pre), a_pre), b_pre); // broadcast scalar + [S]
|
||||
h_pre = ggml_sigmoid(ctx0, h_pre); // [S,T] (+eps omitted, inert)
|
||||
|
||||
// combine: x[h,t] = sum_s h_pre[s,t] * R[h,s,t]
|
||||
//// combine: x[h,t] = sum_s h_pre[s,t] * R[h,s,t]
|
||||
ggml_tensor * hpre3 = ggml_reshape_3d(ctx0, h_pre, 1, S, n_tokens);
|
||||
ggml_tensor * weighted = ggml_mul(ctx0, Rin, hpre3); // [H,S,T]
|
||||
ggml_tensor * x = ggml_reshape_2d(ctx0, ggml_sum_rows_ext(ctx0, weighted, 1), n_embd, n_tokens);
|
||||
auto x = ggml_mul_multi_add(ctx0, Rin, hpre3);
|
||||
//ggml_tensor * weighted = ggml_mul(ctx0, Rin, hpre3); // [H,S,T]
|
||||
//ggml_tensor * x = ggml_reshape_2d(ctx0, ggml_sum_rows_ext(ctx0, weighted, 1), n_embd, n_tokens);
|
||||
ggml_build_forward_expand(gf, x);
|
||||
|
||||
*h_post_out = ggml_cont(ctx0, h_post);
|
||||
|
||||
@@ -549,6 +549,13 @@ ggml_tensor * llm_build_context::build_inp_KQ_mask_swa_win(int64_t n_kv_win, boo
|
||||
return flash_attn ? ggml_cast(ctx0, lctx.inp_KQ_mask_swa_win, GGML_TYPE_F16) : lctx.inp_KQ_mask_swa_win;
|
||||
}
|
||||
|
||||
//build_mhc_post: x = 4096 x 4096 x 1 x 1, post = 4 x 4096 x 1 x 1, residual = 4096 x 4 x 4096 x 1, comb = 4 x 4 x 4096 x 1
|
||||
//build_mhc_post: x = 4096 x 1 x 1 x 1, post = 4 x 1 x 1 x 1, residual = 4096 x 4 x 1 x 1, comb = 4 x 4 x 1 x 1
|
||||
// x = n_embd x n_tokens <--- y in Pangu
|
||||
// post = 4 x n_tokens <--- h_post in Pangu
|
||||
// residual = n_embd x 4 x n_tokens <--- Rin in Pangu
|
||||
// comb = 4 x 4 x n_tokens
|
||||
|
||||
ggml_tensor * llm_build_context::build_mhc_post(
|
||||
ggml_tensor * x,
|
||||
ggml_tensor * post,
|
||||
@@ -560,6 +567,26 @@ ggml_tensor * llm_build_context::build_mhc_post(
|
||||
const int64_t n_tokens = x->ne[1];
|
||||
ggml_tensor * out = nullptr;
|
||||
|
||||
//ggml_tensor * x3 = ggml_reshape_3d(ctx0, x, n_embd, 1, n_tokens);
|
||||
//ggml_tensor * post3 = ggml_reshape_3d(ctx0, post, 1, n_stream, n_tokens);
|
||||
//ggml_tensor repeater;
|
||||
//repeater.ne[0] = n_embd; repeater.ne[1] = n_stream; repeater.ne[2] = n_tokens; repeater.ne[3] = 1;
|
||||
//ggml_tensor * term1 = ggml_mul(ctx0, ggml_repeat(ctx0, x3, &repeater), post3);
|
||||
|
||||
//ggml_tensor * term2 = nullptr;
|
||||
//for (int s = 0; s < n_stream; ++s) {
|
||||
// ggml_tensor * m_s = ggml_cont(ctx0, ggml_view_2d(ctx0, comb, n_stream, n_tokens, comb->nb[2], s*comb->nb[1]));
|
||||
// ggml_tensor * m_s3 = ggml_reshape_3d(ctx0, m_s, 1, n_stream, n_tokens);
|
||||
// ggml_tensor * acc = ggml_mul(ctx0, residual, m_s3);
|
||||
// ggml_tensor * summed = ggml_sum_rows_ext(ctx0, acc, 1);
|
||||
// term2 = term2 ? ggml_concat(ctx0, term2, summed, 1) : summed;
|
||||
//}
|
||||
//return ggml_add(ctx0, term1, term2);
|
||||
//printf("%s: x = %ld x %ld x %ld x %ld, post = %ld x %ld x %ld x %ld, residual = %ld x %ld x %ld x %ld, comb = %ld x %ld x %ld x %ld\n",
|
||||
// __func__, x->ne[0], x->ne[1], x->ne[2], x->ne[3], post->ne[0], post->ne[1], post->ne[2], post->ne[3],
|
||||
// residual->ne[0], residual->ne[1], residual->ne[2], residual->ne[3],
|
||||
// comb->ne[0], comb->ne[1], comb->ne[2], comb->ne[3]);
|
||||
|
||||
for (int64_t dst = 0; dst < n_stream; ++dst) {
|
||||
ggml_tensor * post_dst = ggml_cont(ctx0,
|
||||
ggml_view_2d(ctx0, post, 1, n_tokens, post->nb[1], dst*post->nb[0]));
|
||||
@@ -578,7 +605,13 @@ ggml_tensor * llm_build_context::build_mhc_post(
|
||||
}
|
||||
|
||||
cur = ggml_reshape_3d(ctx0, cur, n_embd, 1, n_tokens);
|
||||
out = out ? ggml_concat(ctx0, out, cur, 1) : cur;
|
||||
if (out) {
|
||||
out = ggml_concat(ctx0, out, cur, 1);
|
||||
cb(out, "mhc_post", dst);
|
||||
} else {
|
||||
out = cur;
|
||||
}
|
||||
//out = out ? ggml_concat(ctx0, out, cur, 1) : cur;
|
||||
}
|
||||
|
||||
return out;
|
||||
|
||||
Reference in New Issue
Block a user