This commit is contained in:
Kawrakow
2026-07-18 06:27:44 +00:00
parent 80d680741f
commit 5ca6285d63
3 changed files with 64 additions and 17 deletions
+26 -13
View File
@@ -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;
+4 -3
View File
@@ -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);
+34 -1
View File
@@ -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;