DS4 refactoring (#2190)

* DS4 refactoring

* Minor
This commit is contained in:
Kawrakow 2026-07-27 09:14:35 +03:00 committed by GitHub
parent 1a7691fae7
commit 5f063b7bba
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 97 additions and 176 deletions

View File

@ -262,7 +262,6 @@ void ggml_cuda_op_concat(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
auto row_size_dst = ggml_row_size(dst->type, dst->ne[0]);
if (row_size_src0 % sizeof(float) == 0 && row_size_src1 % sizeof(float) == 0 && row_size_dst % sizeof(float) == 0) {
auto ne00_eff = row_size_src0/sizeof(float);
auto ne10_eff = row_size_src1/sizeof(float);
auto ne0_eff = row_size_dst /sizeof(float);
dim3 grid(dst->ne[1], dst->ne[2], dst->ne[3]);
k_concat_dim0<<<grid, CUDA_CONCAT_BLOCK_SIZE, 0, ctx.stream()>>>(ne0_eff, ne00_eff,

View File

@ -680,39 +680,6 @@ static ggml_tensor * build_hc_pre(
auto post = ggml_view_2d(ctx0, all, hc, nt, hc*sizeof(float), hc*nt*sizeof(float));
auto comb = ggml_view_3d(ctx0, all, hc, hc, nt, hc*sizeof(float), hc*hc*sizeof(float), 2*hc*nt*sizeof(float));
////ggml_tensor * mixes = llm.build_mhc_pre_projection(x, hc_fn, nullptr,
//// n_embd, hc, norm_rms_eps, false);
////printf("hc_scale: %ld x %ld x %ld x %ld, hc_base: %ld x %ld x %ld x %ld\n",
//// hc_scale->ne[0], hc_scale->ne[1], hc_scale->ne[2], hc_scale->ne[3],
//// hc_base->ne[0], hc_base->ne[1], hc_base->ne[2], hc_base->ne[3]);
//ggml_tensor * scale_pre = dsv4_view_1d(ctx0, hc_scale, 1, 0);
//ggml_tensor * scale_post = dsv4_view_1d(ctx0, hc_scale, 1, 1);
//ggml_tensor * scale_comb = dsv4_view_1d(ctx0, hc_scale, 1, 2);
//ggml_tensor * base_pre = dsv4_view_1d(ctx0, hc_base, hc, 0);
//ggml_tensor * base_post = dsv4_view_1d(ctx0, hc_base, hc, hc);
//ggml_tensor * base_comb = dsv4_view_1d(ctx0, hc_base, hc*hc, 2*hc);
//ggml_tensor * pre = ggml_cont(ctx0, dsv4_view_2d(ctx0, mixes, hc, nt, 0));
//pre = dsv4_hc_affine(ctx0, pre, scale_pre, base_pre);
//pre = ggml_sigmoid(ctx0, pre);
//pre = ggml_scale_bias(ctx0, pre, 1.0f, hparams.dsv4_hc_eps);
//auto post = ggml_cont(ctx0, dsv4_view_2d(ctx0, mixes, hc, nt, hc));
//post = dsv4_hc_affine(ctx0, post, scale_post, base_post);
//post = ggml_sigmoid(ctx0, post);
//post = ggml_scale(ctx0, post, 2.0f);
//auto comb = ggml_cont(ctx0, dsv4_view_2d(ctx0, mixes, hc*hc, nt, 2*hc));
//comb = dsv4_hc_affine(ctx0, comb, scale_comb, base_comb);
//comb = ggml_sinkhorn(ctx0, comb, hc, hparams.dsv4_hc_sinkhorn_iters, hparams.dsv4_hc_eps, false);
////*comb = ggml_reshape_3d(ctx0, *comb, hc, hc, nt);
////*comb = build_hc_sinkhorn(ctx0, hparams, *comb);
//printf("pre: %ld x %ld x %ld x %ld, post: %ld x %ld x %ld x %ld, comb: %ld x %ld x %ld x %ld\n",
// pre->ne[0], pre->ne[1], pre->ne[2], pre->ne[3], post->ne[0], post->ne[1], post->ne[2], post->ne[3],
// comb->ne[0], comb->ne[1], comb->ne[2], comb->ne[3]);
*post_out = post;
*comb_out = comb;
@ -1003,6 +970,75 @@ static ggml_tensor * dsv4_build_lid_top_k(
return top_k;
}
static void ds4_build_comp(ggml_tensor * cur, llm_build_context & llm, ggml_context * ctx0,
llama_context::dsv4_runtime::comp_inputs & inputs,
llama_context::dsv4_runtime::comp_plan & plan,
ggml_tensor * comp_wkv, ggml_tensor * comp_wgate, ggml_tensor * comp_ape, ggml_tensor * norm,
ggml_tensor * cache_state, ggml_tensor * cache_score, ggml_tensor * cache_k,
ggml_tensor ** append_state, ggml_tensor ** append_score,
int head_size, int il, bool do_hadamard, const std::string & tag, ggml_cgraph * gf, bool is_hca) {
ggml_tensor * state_kv = llm.llm_build_lora_mm(llm.lctx, ctx0, comp_wkv, cur);
llm.cb(state_kv, (tag + "_state_kv").c_str(), il);
ggml_tensor * state_score = llm.llm_build_lora_mm(llm.lctx, ctx0, comp_wgate, cur);
llm.cb(state_score, (tag + "_state_score").c_str(), il);
ggml_tensor * ape_rows = ggml_get_rows(ctx0, comp_ape, inputs.state_pos);
llm.cb(ape_rows, (tag + "_ape").c_str(), il);
state_score = ggml_add(ctx0, state_score, ape_rows);
ggml_tensor * dep = nullptr;
if (append_state) {
state_kv = dsv4_append_zero_row(ctx0, state_kv, append_state, false);
}
if (append_score) {
state_score = dsv4_append_zero_row(ctx0, state_score, append_score, true);
}
if (inputs.state_write_idxs != nullptr && plan.state_write_idxs.size() > 0) {
ggml_tensor * source_kv = dsv4_concat_named(ctx0, cache_state, state_kv, 1, (tag + "_source_kv").c_str());
ggml_tensor * source_score = dsv4_concat_named(ctx0, cache_score, state_score, 1, (tag + "_source_score").c_str());
ggml_tensor * comp = !is_hca ? build_overlap_compressed_kv_from_state(ctx0, llm,
source_kv, source_score,
inputs.state_read_idxs,
inputs.state_write_pos,
norm,
llama_context::dsv4_runtime::CSA_RATIO,
head_size, il,
(tag + "_state_compress").c_str()) :
build_hca_compressed_kv_from_state(ctx0, llm,
source_kv, source_score,
inputs.state_read_idxs,
inputs.state_write_pos,
norm, head_size, il);
if (do_hadamard) {
const int hadamard_block = llama_model::hadamard_size(head_size);
GGML_ASSERT(hadamard_block > 0);
GGML_ASSERT(comp->ne[0] % hadamard_block == 0);
comp = ggml_hadamard(ctx0, comp, hadamard_block);
llm.cb(comp, (tag + "_state_compress_hadamard").c_str(), il);
}
ggml_tensor * comp_2d = ggml_reshape_2d(ctx0, comp, head_size, inputs.state_write_idxs->ne[0]);
ggml_tensor * write = dsv4_comp_cpy_k(ctx0, cache_k, comp_2d, inputs.state_write_idxs, head_size);
ggml_build_forward_expand(gf, write);
llm.cb(write, (tag + "_k_write").c_str(), il);
dep = comp;
}
if (dep) {
ggml_build_forward_expand(gf, dep);
}
ggml_tensor * persist_kv = ggml_get_rows(ctx0, state_kv, inputs.state_persist_src_idxs);
llm.cb(persist_kv, (tag + "_persist_kv").c_str(), il);
ggml_tensor * persist_score = ggml_get_rows(ctx0, state_score, inputs.state_persist_src_idxs);
llm.cb(persist_score, (tag + "_persist_score").c_str(), il);
ggml_tensor * state_kv_write = dsv4_comp_state_cpy(ctx0, cache_state, persist_kv, inputs.state_persist_dst_idxs);
ggml_tensor * state_score_write = dsv4_comp_state_cpy(ctx0, cache_score, persist_score, inputs.state_persist_dst_idxs);
ggml_build_forward_expand(gf, state_kv_write);
ggml_build_forward_expand(gf, state_score_write);
llm.cb(state_kv_write, (tag + "_k_state_persist").c_str(), il);
llm.cb(state_score_write, (tag + "_score_state_persist").c_str(), il);
}
ggml_cgraph * llm_build_context::build_deepseek4() {
ggml_cgraph * gf = new_graph_custom();
@ -1122,152 +1158,38 @@ ggml_cgraph * llm_build_context::build_deepseek4() {
}
const float kq_scale = 1.0f / std::sqrt(float(n_embd_head));
ggml_tensor * hca_state_kv = nullptr;
ggml_tensor * hca_state_score = nullptr;
if (ratio == llama_context::dsv4_runtime::HCA_RATIO && lctx.dsv4.inputs.hca.state_pos != nullptr && lctx.dsv4.hca_plan.state_pos.size() > 0) {
hca_state_kv = llm_build_lora_mm(lctx, ctx0, model.layers[il].attn_comp_wkv, cur);
cb(hca_state_kv, "hca_state_kv", il);
hca_state_score = llm_build_lora_mm(lctx, ctx0, model.layers[il].attn_comp_wgate, cur);
cb(hca_state_score, "hca_state_score", il);
ggml_tensor * ape_rows = ggml_get_rows(ctx0, model.layers[il].attn_comp_ape, lctx.dsv4.inputs.hca.state_pos);
cb(ape_rows, "ape", il);
hca_state_score = ggml_add(ctx0, hca_state_score, ape_rows);
cb(hca_state_kv, "hca_state_kv", il);
cb(hca_state_score, "hca_state_score", il);
if (ratio == llama_context::dsv4_runtime::CSA_RATIO &&
lctx.dsv4.inputs.csa.state_pos != nullptr &&
lctx.dsv4.csa_plan.state_pos.size() > 0) {
ds4_build_comp(cur, *this, ctx0, lctx.dsv4.inputs.csa, lctx.dsv4.csa_plan,
model.layers[il].attn_comp_wkv, model.layers[il].attn_comp_wgate,
model.layers[il].attn_comp_ape, model.layers[il].attn_comp_norm,
lctx.dsv4.cache.csa_state_kv[il], lctx.dsv4.cache.csa_state_score[il], lctx.dsv4.cache.csa_k[il],
&append_csa_state, &append_csa_score,
n_embd_head, il, false, "csa", gf, false);
ds4_build_comp(cur, *this, ctx0, lctx.dsv4.inputs.lid, lctx.dsv4.lid_plan,
model.layers[il].indexer_comp_wkv, model.layers[il].indexer_comp_wgate,
model.layers[il].indexer_comp_ape, model.layers[il].indexer_comp_norm,
lctx.dsv4.cache.lid_state_kv[il], lctx.dsv4.cache.lid_state_score[il], lctx.dsv4.cache.lid_k[il],
&append_lid_state, &append_lid_score,
hparams.indexer_head_size, il, true, "lid", gf, false);
}
if (ratio == llama_context::dsv4_runtime::CSA_RATIO && lctx.dsv4.inputs.csa.state_pos != nullptr && lctx.dsv4.csa_plan.state_pos.size() > 0) {
ggml_tensor * csa_state_kv = llm_build_lora_mm(lctx, ctx0, model.layers[il].attn_comp_wkv, cur);
cb(csa_state_kv, "csa_state_kv", il);
ggml_tensor * csa_state_score = llm_build_lora_mm(lctx, ctx0, model.layers[il].attn_comp_wgate, cur);
cb(csa_state_score, "csa_state_score", il);
ggml_tensor * csa_ape_rows = ggml_get_rows(ctx0, model.layers[il].attn_comp_ape, lctx.dsv4.inputs.csa.state_pos);
cb(csa_ape_rows, "csa_ape", il);
csa_state_score = ggml_add(ctx0, csa_state_score, csa_ape_rows);
ggml_tensor * csa_dep = nullptr;
if (ratio == llama_context::dsv4_runtime::HCA_RATIO &&
lctx.dsv4.inputs.hca.state_pos != nullptr &&
lctx.dsv4.hca_plan.state_pos.size() > 0) {
csa_state_kv = dsv4_append_zero_row(ctx0, csa_state_kv, &append_csa_state, false);
csa_state_score = dsv4_append_zero_row(ctx0, csa_state_score, &append_csa_score, true);
ds4_build_comp(cur, *this, ctx0, lctx.dsv4.inputs.hca, lctx.dsv4.hca_plan,
model.layers[il].attn_comp_wkv, model.layers[il].attn_comp_wgate,
model.layers[il].attn_comp_ape, model.layers[il].attn_comp_norm,
lctx.dsv4.cache.hca_state_kv[il], lctx.dsv4.cache.hca_state_score[il], lctx.dsv4.cache.hca_k[il],
nullptr, nullptr,
n_embd_head, il, false, "hca", gf, true);
if (lctx.dsv4.inputs.csa.state_write_idxs != nullptr && lctx.dsv4.csa_plan.state_write_idxs.size() > 0) {
ggml_tensor * csa_source_kv = dsv4_concat_named(ctx0, lctx.dsv4.cache.csa_state_kv[il], csa_state_kv, 1, "dsv4_csa_source_kv");
ggml_tensor * csa_source_score = dsv4_concat_named(ctx0, lctx.dsv4.cache.csa_state_score[il], csa_state_score, 1, "dsv4_csa_source_score");
ggml_tensor * csa_comp = build_overlap_compressed_kv_from_state(
ctx0, *this,
csa_source_kv, csa_source_score,
lctx.dsv4.inputs.csa.state_read_idxs,
lctx.dsv4.inputs.csa.state_write_pos,
model.layers[il].attn_comp_norm,
llama_context::dsv4_runtime::CSA_RATIO,
n_embd_head,
il,
"csa_state_compress");
ggml_tensor * csa_comp_2d = ggml_reshape_2d(ctx0, csa_comp, n_embd_head, lctx.dsv4.inputs.csa.state_write_idxs->ne[0]);
ggml_tensor * csa_write = dsv4_comp_cpy_k(ctx0, lctx.dsv4.cache.csa_k[il], csa_comp_2d, lctx.dsv4.inputs.csa.state_write_idxs, n_embd_head);
ggml_build_forward_expand(gf, csa_write);
cb(csa_write, "dsv4_csa_k_write", il);
csa_dep = csa_comp;
}
if (csa_dep) {
ggml_build_forward_expand(gf, csa_dep);
}
ggml_tensor * csa_persist_kv = ggml_get_rows(ctx0, csa_state_kv, lctx.dsv4.inputs.csa.state_persist_src_idxs);
cb(csa_persist_kv, "csa_persist_kv", il);
ggml_tensor * csa_persist_score = ggml_get_rows(ctx0, csa_state_score, lctx.dsv4.inputs.csa.state_persist_src_idxs);
cb(csa_persist_score, "csa_persist_score", il);
ggml_tensor * csa_state_kv_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.csa_state_kv[il], csa_persist_kv, lctx.dsv4.inputs.csa.state_persist_dst_idxs);
ggml_tensor * csa_state_score_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.csa_state_score[il], csa_persist_score, lctx.dsv4.inputs.csa.state_persist_dst_idxs);
ggml_build_forward_expand(gf, csa_state_kv_write);
ggml_build_forward_expand(gf, csa_state_score_write);
cb(csa_state_kv_write, "dsv4_csa_k_state_persist", il);
cb(csa_state_score_write, "dsv4_csa_score_state_persist", il);
ggml_tensor * lid_state_kv = llm_build_lora_mm(lctx, ctx0, model.layers[il].indexer_comp_wkv, cur);
cb(lid_state_kv, "lid_state_kv", il);
ggml_tensor * lid_state_score = llm_build_lora_mm(lctx, ctx0, model.layers[il].indexer_comp_wgate, cur);
cb(lid_state_score, "lid_state_score", il);
ggml_tensor * lid_ape_rows = ggml_get_rows(ctx0, model.layers[il].indexer_comp_ape, lctx.dsv4.inputs.lid.state_pos);
cb(lid_ape_rows, "lid_ape", il);
lid_state_score = ggml_add(ctx0, lid_state_score, lid_ape_rows);
ggml_tensor * lid_dep = nullptr;
lid_state_kv = dsv4_append_zero_row(ctx0, lid_state_kv, &append_lid_state, false);
lid_state_score = dsv4_append_zero_row(ctx0, lid_state_score, &append_lid_score, true);
if (lctx.dsv4.inputs.lid.state_write_idxs != nullptr && lctx.dsv4.lid_plan.state_write_idxs.size() > 0) {
ggml_tensor * lid_source_kv = dsv4_concat_named(ctx0, lctx.dsv4.cache.lid_state_kv[il], lid_state_kv, 1, "dsv4_lid_source_kv");
ggml_tensor * lid_source_score = dsv4_concat_named(ctx0, lctx.dsv4.cache.lid_state_score[il], lid_state_score, 1, "dsv4_lid_source_score");
ggml_tensor * lid_comp = build_overlap_compressed_kv_from_state(
ctx0, *this,
lid_source_kv, lid_source_score,
lctx.dsv4.inputs.lid.state_read_idxs,
lctx.dsv4.inputs.lid.state_write_pos,
model.layers[il].indexer_comp_norm,
llama_context::dsv4_runtime::CSA_RATIO,
hparams.indexer_head_size,
il,
"lid_state_compress");
const int hadamard_block = llama_model::hadamard_size((int) hparams.indexer_head_size);
GGML_ASSERT(hadamard_block > 0);
GGML_ASSERT(lid_comp->ne[0] % hadamard_block == 0);
lid_comp = ggml_hadamard(ctx0, lid_comp, hadamard_block);
cb(lid_comp, "lid_state_compress_hadamard", il);
ggml_tensor * lid_comp_2d = ggml_reshape_2d(ctx0, lid_comp, hparams.indexer_head_size, lctx.dsv4.inputs.lid.state_write_idxs->ne[0]);
ggml_tensor * lid_write = dsv4_comp_cpy_k(ctx0, lctx.dsv4.cache.lid_k[il], lid_comp_2d, lctx.dsv4.inputs.lid.state_write_idxs, hparams.indexer_head_size);
ggml_build_forward_expand(gf, lid_write);
cb(lid_write, "dsv4_lid_k_write", il);
lid_dep = lid_comp;
}
if (lid_dep) {
ggml_build_forward_expand(gf, lid_dep);
}
ggml_tensor * lid_persist_kv = ggml_get_rows(ctx0, lid_state_kv, lctx.dsv4.inputs.lid.state_persist_src_idxs);
cb(lid_persist_kv, "lid_persist_kv", il);
ggml_tensor * lid_persist_score = ggml_get_rows(ctx0, lid_state_score, lctx.dsv4.inputs.lid.state_persist_src_idxs);
cb(lid_persist_score, "lid_persist_score", il);
ggml_tensor * lid_state_kv_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.lid_state_kv[il], lid_persist_kv, lctx.dsv4.inputs.lid.state_persist_dst_idxs);
ggml_tensor * lid_state_score_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.lid_state_score[il], lid_persist_score, lctx.dsv4.inputs.lid.state_persist_dst_idxs);
ggml_build_forward_expand(gf, lid_state_kv_write);
ggml_build_forward_expand(gf, lid_state_score_write);
cb(lid_state_kv_write, "dsv4_lid_k_state_persist", il);
cb(lid_state_score_write, "dsv4_lid_score_state_persist", il);
}
if (ratio == llama_context::dsv4_runtime::HCA_RATIO && hca_state_kv != nullptr && hca_state_score != nullptr) {
ggml_tensor * hca_dep = nullptr;
if (lctx.dsv4.inputs.hca.state_write_idxs != nullptr && lctx.dsv4.hca_plan.state_write_idxs.size() > 0) {
ggml_tensor * hca_source_kv = dsv4_concat_named(ctx0, lctx.dsv4.cache.hca_state_kv[il], hca_state_kv, 1, "dsv4_hca_source_kv");
ggml_tensor * hca_source_score = dsv4_concat_named(ctx0, lctx.dsv4.cache.hca_state_score[il], hca_state_score, 1, "dsv4_hca_source_score");
ggml_tensor * hca_comp = build_hca_compressed_kv_from_state(
ctx0, *this,
hca_source_kv, hca_source_score,
lctx.dsv4.inputs.hca.state_read_idxs,
lctx.dsv4.inputs.hca.state_write_pos,
model.layers[il].attn_comp_norm,
n_embd_head,
il);
ggml_tensor * hca_comp_2d = ggml_reshape_2d(ctx0, hca_comp, n_embd_head, lctx.dsv4.inputs.hca.state_write_idxs->ne[0]);
ggml_tensor * hca_write = dsv4_comp_cpy_k(ctx0, lctx.dsv4.cache.hca_k[il], hca_comp_2d, lctx.dsv4.inputs.hca.state_write_idxs, n_embd_head);
ggml_build_forward_expand(gf, hca_write);
cb(hca_write, "dsv4_hca_k_write", il);
hca_dep = hca_comp;
}
if (hca_dep) {
ggml_build_forward_expand(gf, hca_dep);
}
ggml_tensor * hca_persist_kv = ggml_get_rows(ctx0, hca_state_kv, lctx.dsv4.inputs.hca.state_persist_src_idxs);
ggml_tensor * hca_persist_score = ggml_get_rows(ctx0, hca_state_score, lctx.dsv4.inputs.hca.state_persist_src_idxs);
cb(hca_persist_kv, "hca_persist_kv", il);
cb(hca_persist_score, "hca_persist_score", il);
ggml_tensor * hca_state_kv_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.hca_state_kv[il], hca_persist_kv, lctx.dsv4.inputs.hca.state_persist_dst_idxs);
ggml_tensor * hca_state_score_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.hca_state_score[il], hca_persist_score, lctx.dsv4.inputs.hca.state_persist_dst_idxs);
ggml_build_forward_expand(gf, hca_state_kv_write);
ggml_build_forward_expand(gf, hca_state_score_write);
cb(hca_state_kv_write, "dsv4_hca_k_state_persist", il);
cb(hca_state_score_write, "dsv4_hca_score_state_persist", il);
}
ggml_tensor * raw_k_write = nullptr;