Remove DSV4 zero-dependency graph shim

This commit is contained in:
samuel
2026-07-14 12:10:01 -03:00
parent 724ebba103
commit 050021d5f3
+15 -30
View File
@@ -102,23 +102,6 @@ static ggml_tensor * dsv4_append_zero_row(ggml_context * ctx, ggml_tensor * t, b
return dsv4_concat_named(ctx, t, row, 1, "dsv4_append_zero_row");
}
static ggml_tensor * dsv4_with_zero_dep(ggml_context * ctx, ggml_tensor * t, ggml_tensor * dep) {
// Consumers of t must depend on dep so compression/current-cache work completes
// before state persistence or dependent cache reads, graph expansion order alone
// is not a scheduler data dependency. Keep the typed scalar/add1 form because
// the F16 raw-K path requires it.
if (dep == nullptr) {
return t;
}
ggml_tensor * zero = ggml_sum(ctx, dep);
if (zero->type != GGML_TYPE_F32) {
zero = ggml_cast(ctx, zero, GGML_TYPE_F32);
}
zero = ggml_scale(ctx, zero, 0.0f);
return ggml_add1(ctx, t, zero);
}
static ggml_tensor * dsv4_cache_view_2d(
ggml_context * ctx,
ggml_tensor * cache,
@@ -259,8 +242,7 @@ static ggml_tensor * dsv4_raw_get_k(
ggml_context * ctx,
ggml_tensor * cache,
ggml_tensor * raw_k_read_idxs,
int64_t n_embd_head,
ggml_tensor * dep) {
int64_t n_embd_head) {
if (cache == nullptr) {
return nullptr;
}
@@ -300,7 +282,6 @@ static ggml_tensor * dsv4_raw_get_k(
ggml_tensor * raw_k = ggml_reshape_4d(ctx, rows, n_embd_head, n_head_kv, n_kv, n_stream);
GGML_UNUSED(dep);
return raw_k;
}
@@ -1025,6 +1006,7 @@ ggml_cgraph * llm_build_context::build_deepseek4() {
ggml_tensor * csa_state_score = llm_build_lora_mm(lctx, ctx0, model.layers[il].attn_comp_wgate, cur);
ggml_tensor * csa_ape_rows = ggml_get_rows(ctx0, model.layers[il].attn_comp_ape, lctx.dsv4.inputs.csa.state_pos);
csa_state_score = ggml_add(ctx0, csa_state_score, csa_ape_rows);
ggml_tensor * csa_dep = nullptr;
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[(size_t) il], csa_state_kv, 1, "dsv4_csa_source_kv");
@@ -1043,10 +1025,12 @@ ggml_cgraph * llm_build_context::build_deepseek4() {
ggml_tensor * csa_write = dsv4_comp_cpy_k(ctx0, lctx.dsv4.cache.csa_k[(size_t) 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_state_kv = dsv4_with_zero_dep(ctx0, csa_state_kv, csa_comp);
csa_state_score = dsv4_with_zero_dep(ctx0, csa_state_score, csa_comp);
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);
ggml_tensor * csa_persist_score = ggml_get_rows(ctx0, csa_state_score, lctx.dsv4.inputs.csa.state_persist_src_idxs);
ggml_tensor * csa_state_kv_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.csa_state_kv[(size_t) il], csa_persist_kv, lctx.dsv4.inputs.csa.state_persist_dst_idxs);
@@ -1060,6 +1044,7 @@ ggml_cgraph * llm_build_context::build_deepseek4() {
ggml_tensor * lid_state_score = llm_build_lora_mm(lctx, ctx0, model.layers[il].indexer_comp_wgate, cur);
ggml_tensor * lid_ape_rows = ggml_get_rows(ctx0, model.layers[il].indexer_comp_ape, lctx.dsv4.inputs.lid.state_pos);
lid_state_score = ggml_add(ctx0, lid_state_score, lid_ape_rows);
ggml_tensor * lid_dep = nullptr;
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[(size_t) il], lid_state_kv, 1, "dsv4_lid_source_kv");
@@ -1083,10 +1068,12 @@ ggml_cgraph * llm_build_context::build_deepseek4() {
ggml_tensor * lid_write = dsv4_comp_cpy_k(ctx0, lctx.dsv4.cache.lid_k[(size_t) 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_state_kv = dsv4_with_zero_dep(ctx0, lid_state_kv, lid_comp);
lid_state_score = dsv4_with_zero_dep(ctx0, lid_state_score, lid_comp);
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);
ggml_tensor * lid_persist_score = ggml_get_rows(ctx0, lid_state_score, lctx.dsv4.inputs.lid.state_persist_src_idxs);
ggml_tensor * lid_state_kv_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.lid_state_kv[(size_t) il], lid_persist_kv, lctx.dsv4.inputs.lid.state_persist_dst_idxs);
@@ -1117,9 +1104,9 @@ ggml_cgraph * llm_build_context::build_deepseek4() {
hca_dep = hca_comp;
}
hca_state_kv = dsv4_with_zero_dep(ctx0, hca_state_kv, hca_dep);
hca_state_score = dsv4_with_zero_dep(ctx0, hca_state_score, hca_dep);
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);
ggml_tensor * hca_state_kv_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.hca_state_kv[(size_t) il], hca_persist_kv, lctx.dsv4.inputs.hca.state_persist_dst_idxs);
@@ -1143,9 +1130,8 @@ ggml_cgraph * llm_build_context::build_deepseek4() {
llm_build_kv_store(lctx, ctx0, hparams, cparams, kv_self, gf, nullptr, kv, n_tokens, kv_head, cb, il);
ggml_tensor * raw_k = nullptr;
ggml_tensor * raw_k_dep = 2*il < (int64_t) lctx.cache_copies.size() ? lctx.cache_copies[2*il + 0].cpy : nullptr;
if (hparams.n_head_kv(il) == 1 && lctx.dsv4.inputs.raw_k_read_idxs != nullptr) {
raw_k = dsv4_raw_get_k(&lctx, ctx0, kv_self.k_l[il], lctx.dsv4.inputs.raw_k_read_idxs, n_embd_head, raw_k_dep);
raw_k = dsv4_raw_get_k(&lctx, ctx0, kv_self.k_l[il], lctx.dsv4.inputs.raw_k_read_idxs, n_embd_head);
}
if (raw_k == nullptr) {
raw_k = ggml_view_3d(ctx0, kv_self.k_l[il],
@@ -1153,7 +1139,6 @@ ggml_cgraph * llm_build_context::build_deepseek4() {
ggml_row_size(kv_self.k_l[il]->type, n_embd_head),
ggml_row_size(kv_self.k_l[il]->type, n_embd_head) * hparams.n_head_kv(il),
0);
raw_k = dsv4_with_zero_dep(ctx0, raw_k, raw_k_dep);
}
cb(raw_k, "raw_k", il);