Refactor DSV4 tensor handling for MTP execution and improve raw context management

This commit is contained in:
samuel
2026-07-12 11:44:48 -03:00
parent 88f5247203
commit 2cf9fd95bc
3 changed files with 120 additions and 28 deletions
+35 -26
View File
@@ -323,8 +323,6 @@ static ggml_tensor * dsv4_raw_get_k(
ggml_tensor * raw_k_read_idxs,
int64_t n_embd_head,
ggml_tensor * dep) {
GGML_UNUSED(raw_k_read_idxs);
if (cache == nullptr) {
return nullptr;
}
@@ -350,17 +348,19 @@ static ggml_tensor * dsv4_raw_get_k(
GGML_ASSERT(n_embd_gqa % n_embd_head == 0);
const int64_t n_head_kv = n_embd_gqa/n_embd_head;
const size_t row_size_head = ggml_row_size(cache->type, n_embd_head);
const size_t row_size_gqa = ggml_row_size(cache->type, n_embd_gqa);
const int64_t kv_size = cache->ne[1];
const size_t offset = row_size_gqa*(size_t) sinfo.s0*(size_t) kv_size;
GGML_ASSERT(raw_k_read_idxs != nullptr);
GGML_ASSERT(raw_k_read_idxs->ne[0] >= n_kv*n_stream);
ggml_tensor * raw_k = ggml_view_4d(ctx, cache,
n_embd_head, n_head_kv, n_kv, n_stream,
row_size_head,
row_size_gqa,
row_size_gqa*(size_t) kv_size,
offset);
// Gather controller-owned slots into the attention layout.
ggml_tensor * cache_2d = dsv4_cache_view_2d(ctx, cache, n_embd_gqa, cache->ne[1]);
ggml_tensor * idxs = raw_k_read_idxs->type == GGML_TYPE_I32
? raw_k_read_idxs : ggml_cast(ctx, raw_k_read_idxs, GGML_TYPE_I32);
ggml_tensor * rows = ggml_get_rows(ctx, cache_2d, idxs);
if (rows->type != cache->type && !ggml_is_quantized(cache->type)) {
rows = ggml_cast(ctx, rows, cache->type);
}
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;
@@ -388,7 +388,7 @@ static ggml_tensor * dsv4_raw_cpy_k(
ggml_tensor * write = nullptr;
const auto & sinfo = lctx->dsv4.raw.sinfo_write;
if (cur_2d->ne[1] == raw_k_write_idxs->ne[0]) {
if (sinfo.n_stream() <= 1 && cur_2d->ne[1] == raw_k_write_idxs->ne[0]) {
cur_2d = dsv4_require_f32_rows(ctx, cur_2d);
write = ggml_set_rows(ctx, cache_2d, cur_2d, raw_k_write_idxs);
} else if (sinfo.n_stream() <= 1) {
@@ -400,13 +400,15 @@ static ggml_tensor * dsv4_raw_cpy_k(
const int64_t n_fanout = (int64_t) sinfo.size()*(int64_t) sinfo.n_stream();
GGML_ASSERT(sinfo.n_stream() > 1);
GGML_ASSERT(cur_2d->ne[1] == (int64_t) sinfo.size());
GGML_ASSERT(raw_k_write_idxs->ne[0] == n_fanout);
GGML_UNUSED(raw_k_write_src_idxs);
GGML_ASSERT(raw_k_write_src_idxs->ne[0] == n_fanout);
for (uint32_t s = 0; s < sinfo.n_stream(); ++s) {
ggml_tensor * src_idxs_s = ggml_view_1d(ctx, raw_k_write_src_idxs, sinfo.size(),
s*sinfo.size()*ggml_element_size(raw_k_write_src_idxs));
ggml_tensor * k_idxs_s = ggml_view_1d(ctx, raw_k_write_idxs, sinfo.size(), s*sinfo.size()*ggml_element_size(raw_k_write_idxs));
ggml_tensor * cur_f32 = dsv4_require_f32_rows(ctx, cur_2d);
ggml_tensor * cur_sel = ggml_get_rows(ctx, cur_2d, src_idxs_s);
ggml_tensor * cur_f32 = dsv4_require_f32_rows(ctx, cur_sel);
ggml_tensor * cur = ggml_set_rows(ctx, cache_2d, cur_f32, k_idxs_s);
if (write == nullptr) {
write = cur;
@@ -586,8 +588,9 @@ static ggml_tensor * build_hc_weighted_sum(
ggml_tensor * acc = nullptr;
for (int64_t ih = 0; ih < hc; ++ih) {
ggml_tensor * xh = ggml_view_2d(ctx0, x, n_embd, nt, x->nb[2], ih*x->nb[1]);
ggml_tensor * wh = ggml_view_2d(ctx0, weights, 1, nt, weights->nb[1], ih*weights->nb[0]);
// Materialize strided slices before broadcast operations.
ggml_tensor * xh = ggml_cont(ctx0, ggml_view_2d(ctx0, x, n_embd, nt, x->nb[2], ih*x->nb[1]));
ggml_tensor * wh = ggml_cont(ctx0, ggml_view_2d(ctx0, weights, 1, nt, weights->nb[1], ih*weights->nb[0]));
ggml_tensor * cur = ggml_mul(ctx0, xh, wh);
acc = acc ? ggml_add(ctx0, acc, cur) : cur;
}
@@ -603,7 +606,6 @@ static ggml_tensor * build_hc_sinkhorn(
ggml_tensor * eps = ggml_new_tensor_1d(ctx0, GGML_TYPE_F32, 1);
eps = ggml_fill(ctx0, eps, hparams.dsv4_hc_eps);
comb = ggml_add(ctx0, comb, eps);
auto norm_cols = [&]() {
@@ -657,17 +659,17 @@ static ggml_tensor * build_hc_pre(
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 = dsv4_view_2d(ctx0, mixes, hc, nt, 0);
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);
*post = dsv4_view_2d(ctx0, mixes, hc, nt, hc);
*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);
*comb = dsv4_view_2d(ctx0, mixes, hc*hc, nt, 2*hc);
*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_reshape_3d(ctx0, *comb, hc, hc, nt);
*comb = build_hc_sinkhorn(ctx0, hparams, *comb);
@@ -688,12 +690,12 @@ static ggml_tensor * build_hc_post(
ggml_tensor * out = nullptr;
for (int64_t dst = 0; dst < hc; ++dst) {
ggml_tensor * post_dst = ggml_view_2d(ctx0, post, 1, nt, post->nb[1], dst*post->nb[0]);
ggml_tensor * post_dst = ggml_cont(ctx0, ggml_view_2d(ctx0, post, 1, nt, post->nb[1], dst*post->nb[0]));
ggml_tensor * cur = ggml_mul(ctx0, x, post_dst);
for (int64_t src = 0; src < hc; ++src) {
ggml_tensor * res_src = ggml_view_2d(ctx0, residual, n_embd, nt, residual->nb[2], src*residual->nb[1]);
ggml_tensor * comb_src_dst = ggml_view_2d(ctx0, comb, 1, nt, comb->nb[2], dst*comb->nb[0] + src*comb->nb[1]);
ggml_tensor * res_src = ggml_cont(ctx0, ggml_view_2d(ctx0, residual, n_embd, nt, residual->nb[2], src*residual->nb[1]));
ggml_tensor * comb_src_dst = ggml_cont(ctx0, ggml_view_2d(ctx0, comb, 1, nt, comb->nb[2], dst*comb->nb[0] + src*comb->nb[1]));
cur = ggml_add(ctx0, cur, ggml_mul(ctx0, res_src, comb_src_dst));
}
@@ -966,6 +968,13 @@ static ggml_tensor * dsv4_build_lid_top_k(
ggml_cgraph * llm_build_context::build_deepseek4() {
ggml_cgraph * gf = new_graph_custom();
if (lctx.cparams.mtp_op_type != MTP_OP_NONE) {
GGML_ABORT("DeepSeek4 MTP execution is not implemented");
}
// Exclude the optional NextN tail from ordinary generation.
const int64_t n_base_layers = std::max<int64_t>(0, n_layer - hparams.nextn_predict_layers);
const int64_t n_embd_head = hparams.n_embd_head_k(0);
const int64_t n_embd_head_rope = hparams.n_rot;
const int64_t n_embd_head_nope = n_embd_head - n_embd_head_rope;
@@ -988,7 +997,7 @@ ggml_cgraph * llm_build_context::build_deepseek4() {
inpL = ggml_repeat_4d(ctx0, inpL, n_embd, hc, n_tokens, 1);
cb(inpL, "hc_init", -1);
for (int il = 0; il < n_layer; ++il) {
for (int il = 0; il < n_base_layers; ++il) {
ggml_tensor * residual = inpL;
ggml_tensor * post = nullptr;
ggml_tensor * comb = nullptr;
+49 -1
View File
@@ -373,6 +373,7 @@ static bool dsv4_build_raw_context(
for (size_t s = 0; s < raw.sinfo_read.n_stream(); ++s) {
const llama_seq_id seq_id = read_seq_ids[s];
raw.sinfo_read.idxs[s].clear();
int32_t count = 0;
for (uint32_t slot = 0; slot < kv.size; ++slot) {
const llama_kv_cell & cell = kv.cells[slot];
@@ -405,6 +406,51 @@ static bool dsv4_build_raw_context(
}
}
if (raw.sinfo_write.n_stream() > 1) {
std::vector<int32_t> write_src_idxs;
std::vector<int32_t> write_dst_idxs;
const size_t rows_per_stream = raw.sinfo_write.size();
for (size_t s = 0; s < raw.sinfo_write.n_stream(); ++s) {
if (raw.sinfo_write.idxs[s].size() != rows_per_stream) {
LLAMA_LOG_ERROR("%s: DSV4 packed batch has unequal raw-write rows per stream\n", __func__);
return false;
}
for (int32_t i = 0; i < batch.n_tokens; ++i) {
if (dsv4_token_has_seq(batch, i, write_seq_ids[s])) {
write_src_idxs.push_back(i);
}
}
for (uint32_t slot : raw.sinfo_write.idxs[s]) {
write_dst_idxs.push_back((int32_t) slot);
}
}
raw.write_src_idxs = std::move(write_src_idxs);
raw.write_dst_idxs = std::move(write_dst_idxs);
}
// The graph exposes a rectangular raw-key view. Repeat the last valid row
// for shorter streams; the corresponding mask entries remain -INFINITY.
// This preserves the logical visibility while allowing one get_rows op to
// serve all streams.
if (raw.n_kv > 0) {
raw.read_dst_idxs.clear();
const size_t read_rows = GGML_PAD((size_t) raw.n_kv, 256u);
for (size_t s = 0; s < raw.sinfo_read.n_stream(); ++s) {
const auto & rows = raw.sinfo_read.idxs[s];
for (uint32_t slot : rows) {
raw.read_dst_idxs.push_back((int32_t) slot);
}
const int32_t pad = rows.empty() ? 0 : (int32_t) rows.back();
for (size_t i = rows.size(); i < read_rows; ++i) {
raw.read_dst_idxs.push_back(pad);
}
}
}
return true;
}
@@ -764,7 +810,9 @@ static void dsv4_set_mask_tensor(
}
bool llama_context::ensure_dsv4_cache_tensors() {
const int32_t n_layer = model.hparams.n_layer;
// Allocate cache state only for base-generation layers.
const int32_t n_layer = std::max<int32_t>(
0, model.hparams.n_layer - (int32_t) model.hparams.nextn_predict_layers);
const int64_t n_embd_head = model.hparams.n_embd_head_k(0);
const int64_t n_indexer_head = model.hparams.indexer_head_size;
const uint32_t n_stream = std::max<uint32_t>(1, cparams.n_seq_max);
+36 -1
View File
@@ -2736,8 +2736,12 @@ bool create_tensors_helper::create_deepseek2_tensors(const LLM_TN & tn) {
bool create_tensors_helper::create_deepseek4_tensors(const LLM_TN & tn) {
LOADING_PRELUDE
// Apply the per-layer MTP tail policy to metadata-based tensors.
int layer_flags = 0;
auto create_tensor_from_meta = [&](ggml_context * ctx, const std::string & name, int flags = 0) -> ggml_tensor * {
ggml_tensor * meta = (flags & llama_model_loader::TENSOR_NOT_REQUIRED)
flags |= layer_flags;
ggml_tensor * meta = (flags & (llama_model_loader::TENSOR_NOT_REQUIRED | llama_model_loader::TENSOR_SKIP))
? ml.get_tensor_meta(name.c_str())
: ml.require_tensor_meta(name.c_str());
if (meta == nullptr) {
@@ -2776,6 +2780,15 @@ bool create_tensors_helper::create_deepseek4_tensors(const LLM_TN & tn) {
model.hc_head_scale = create_tensor_from_meta(ctx_output, pick_tensor_name({"hc_head_scale.weight", "output_hc_scale.weight"}), llama_model_loader::TENSOR_NOT_REQUIRED);
for (int i = 0; i < n_layer; ++i) {
const bool is_mtp_layer = hparams.nextn_predict_layers > 0 &&
static_cast<uint32_t>(i) >= n_layer - hparams.nextn_predict_layers;
layer_flags = 0;
if (!model.mtp && is_mtp_layer) {
// Skip the MTP tail when ordinary execution is selected.
layer_flags = llama_model_loader::TENSOR_SKIP | llama_model_loader::TENSOR_NOT_REQUIRED;
}
ggml_context * ctx_split = ctx_for_layer_split(i);
auto & layer = model.layers[i];
@@ -2858,6 +2871,28 @@ bool create_tensors_helper::create_deepseek4_tensors(const LLM_TN & tn) {
format("blk.%d.exp_probs_b.weight", i),
}), llama_model_loader::TENSOR_NOT_REQUIRED);
}
if (is_mtp_layer) {
const int nextn_flags = llama_model_loader::TENSOR_NOT_REQUIRED;
layer.nextn.eh_proj = create_tensor(ctx_split,
tn(LLM_TENSOR_NEXTN_EH_PROJ, "weight", i),
{ 2 * n_embd, n_embd }, layer_flags | nextn_flags);
layer.nextn.embed_tokens = create_tensor(ctx_input,
tn(LLM_TENSOR_NEXTN_EMBED_TOKENS, "weight", i),
{ n_embd, n_vocab }, layer_flags | nextn_flags);
layer.nextn.enorm = create_tensor(ctx_split,
tn(LLM_TENSOR_NEXTN_ENORM, "weight", i),
{ n_embd }, layer_flags | nextn_flags);
layer.nextn.hnorm = create_tensor(ctx_split,
tn(LLM_TENSOR_NEXTN_HNORM, "weight", i),
{ n_embd }, layer_flags | nextn_flags);
layer.nextn.shared_head_head = create_tensor(ctx_output,
tn(LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, "weight", i),
{ n_embd, n_vocab }, layer_flags | nextn_flags);
layer.nextn.shared_head_norm = create_tensor(ctx_output,
tn(LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, "weight", i),
{ n_embd }, layer_flags | nextn_flags);
}
}
return use_mmap_buffer;