DeepseekV4: fix rollback with multi-seq (#26756)
* DeepseekV4: fix rollback with multi-seq * fix model loading * make pending rollback single use * only clear cache for seq_id for full load * add assert for compress ratio * make graph topology static * pass true instead of flags in clear_compressed * cont : clean-up + TODOs --------- Co-authored-by: Georgi Gerganov <ggerganov@gmail.com>
This commit is contained in:
co-authored by
Georgi Gerganov
parent
d3371929bb
commit
b0539c43ed
@@ -228,6 +228,15 @@ if (NOT WIN32 OR NOT BUILD_SHARED_LIBS)
|
||||
set_tests_properties(test-recurrent-state-rollback-nemotron-h PROPERTIES
|
||||
FIXTURES_REQUIRED generate-models
|
||||
)
|
||||
llama_test(
|
||||
test-recurrent-state-rollback
|
||||
NAME test-recurrent-state-rollback-dsv4
|
||||
LABEL main
|
||||
ARGS -m "${MODEL_DIR}/deepseek4-moe.gguf"
|
||||
)
|
||||
set_tests_properties(test-recurrent-state-rollback-dsv4 PROPERTIES
|
||||
FIXTURES_REQUIRED generate-models
|
||||
)
|
||||
endif()
|
||||
|
||||
llama_build_and_test(test-chat-peg-parser.cpp peg-parser/simple-tokenize.cpp)
|
||||
|
||||
@@ -101,6 +101,11 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) {
|
||||
n_head = 1;
|
||||
n_ff = 96;
|
||||
n_layer = 22; // hparams.n_layer_kv_from_start = 20 is hardcoded
|
||||
} else if (arch == LLM_ARCH_DEEPSEEK4) {
|
||||
n_embd = 128;
|
||||
n_head = 1;
|
||||
n_ff = 192;
|
||||
n_layer = 3; // uncompressed + csa + hca, one layer of each ratio kind
|
||||
} else if (arch == LLM_ARCH_STEP35 || arch == LLM_ARCH_LAGUNA) {
|
||||
n_embd = 160; // exercise per-head tensor split granularity with head size 80
|
||||
} else if (arch == LLM_ARCH_QWEN3 || arch == LLM_ARCH_MUSE_GLIMMER || arch == LLM_ARCH_AFMOE) {
|
||||
@@ -203,6 +208,10 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) {
|
||||
}
|
||||
ms.add_kv(LLM_KV_ATTENTION_INDEXER_TYPES, indexer_types);
|
||||
}
|
||||
} else if (arch == LLM_ARCH_DEEPSEEK4) {
|
||||
ms.add_kv(LLM_KV_ATTENTION_KEY_LENGTH, uint32_t(128));
|
||||
ms.add_kv(LLM_KV_ATTENTION_VALUE_LENGTH, uint32_t(128));
|
||||
ms.add_kv(LLM_KV_ROPE_DIMENSION_COUNT, uint32_t(64));
|
||||
} else if (arch == LLM_ARCH_MINIMAX_M3) {
|
||||
// partial rotary: n_rot must not exceed the indexer key length (64)
|
||||
ms.add_kv(LLM_KV_ROPE_DIMENSION_COUNT, uint32_t(64));
|
||||
@@ -239,6 +248,20 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) {
|
||||
|
||||
// MSA requires one indexer head per GQA (KV) head, unlike the DSA archs where the
|
||||
// indexer head count is independent of the main attention head count.
|
||||
if (arch == LLM_ARCH_DEEPSEEK4) {
|
||||
ms.add_kv(LLM_KV_EXPERT_WEIGHTS_SCALE, 2.5f);
|
||||
ms.add_kv(LLM_KV_EXPERT_WEIGHTS_NORM, true);
|
||||
ms.add_kv(LLM_KV_SWIGLU_CLAMP_EXP, 7.0f);
|
||||
ms.add_kv(LLM_KV_ATTENTION_OUTPUT_GROUP_COUNT, uint32_t(1));
|
||||
ms.add_kv(LLM_KV_ATTENTION_OUTPUT_LORA_RANK, uint32_t(64));
|
||||
ms.add_kv(LLM_KV_ATTENTION_COMPRESS_ROPE_FREQ_BASE, 10000.0f);
|
||||
ms.add_kv(LLM_KV_HYPER_CONNECTION_COUNT, uint32_t(4));
|
||||
ms.add_kv(LLM_KV_HYPER_CONNECTION_SINKHORN_ITERATIONS, uint32_t(4));
|
||||
ms.add_kv(LLM_KV_HYPER_CONNECTION_EPSILON, 1e-6f);
|
||||
ms.add_kv(LLM_KV_HASH_LAYER_COUNT, uint32_t(0));
|
||||
ms.add_kv(LLM_KV_ATTENTION_COMPRESS_RATIOS, std::vector<uint32_t>({0, 4, 128}));
|
||||
}
|
||||
|
||||
ms.add_kv(LLM_KV_ATTENTION_INDEXER_HEAD_COUNT, arch == LLM_ARCH_MINIMAX_M3 ? n_head : uint32_t(1));
|
||||
ms.add_kv(LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, uint32_t(64));
|
||||
ms.add_kv(LLM_KV_ATTENTION_INDEXER_TOP_K, uint32_t(8));
|
||||
@@ -257,7 +280,7 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) {
|
||||
ms.add_kv(LLM_KV_EXPERT_COUNT, uint32_t(2));
|
||||
ms.add_kv(LLM_KV_EXPERT_USED_COUNT, uint32_t(1));
|
||||
ms.add_kv(LLM_KV_EXPERT_SHARED_COUNT, uint32_t(1));
|
||||
ms.add_kv(LLM_KV_EXPERT_GATING_FUNC, uint32_t(2)); // sigmoid
|
||||
ms.add_kv(LLM_KV_EXPERT_GATING_FUNC, arch == LLM_ARCH_DEEPSEEK4 ? uint32_t(4) : uint32_t(2)); // sqrtsoftplus : sigmoid
|
||||
ms.add_kv(LLM_KV_EXPERT_GROUP_SCALE, 1.0f);
|
||||
ms.add_kv(LLM_KV_EXPERTS_PER_GROUP, uint32_t(1));
|
||||
}
|
||||
@@ -395,6 +418,7 @@ static bool moe_mandatory(const llm_arch arch) {
|
||||
case LLM_ARCH_DEEPSEEK2:
|
||||
case LLM_ARCH_DEEPSEEK32:
|
||||
case LLM_ARCH_DOTS3NOTE:
|
||||
case LLM_ARCH_DEEPSEEK4:
|
||||
case LLM_ARCH_GLM4_MOE:
|
||||
case LLM_ARCH_GLM_DSA:
|
||||
case LLM_ARCH_EXAONE_MOE:
|
||||
@@ -480,9 +504,6 @@ static bool arch_supported(const llm_arch arch) {
|
||||
if (arch == LLM_ARCH_DEEPSEEK2OCR) {
|
||||
return false;
|
||||
}
|
||||
if (arch == LLM_ARCH_DEEPSEEK4) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// FIXME: these hit scheduler/view-backed-output issues with WebGPU on CI.
|
||||
#ifdef GGML_USE_WEBGPU
|
||||
|
||||
@@ -35,6 +35,178 @@ static bool decode_one(llama_context * ctx, llama_token tok, llama_pos pos) {
|
||||
return ok;
|
||||
}
|
||||
|
||||
// Roll back multiple sequences, then replay them in a single batch whose
|
||||
// per-seq token count exceeds n_ubatch: each seq's replay spans several
|
||||
// ubatches while its rollback restore is still pending. Compared against a
|
||||
// reference context that never advanced past the rollback point and decodes
|
||||
// the identical replay batch.
|
||||
static bool test_multi_seq_split_replay(const common_params & params, llama_model * model, const int n_vocab) {
|
||||
constexpr uint32_t n_seqs = 2;
|
||||
constexpr uint32_t n_ubatch = 16;
|
||||
constexpr uint32_t n_prompt = 19;
|
||||
constexpr uint32_t n_rollback = 3;
|
||||
constexpr uint32_t n_replay = 40; // > n_ubatch so each seq spans multiple ubatches
|
||||
constexpr llama_pos p0 = n_prompt - n_rollback;
|
||||
|
||||
const auto make_ctx_multi = [&]() {
|
||||
auto cparams = common_context_params_to_llama(params);
|
||||
cparams.n_seq_max = n_seqs;
|
||||
cparams.n_rs_seq = 8;
|
||||
cparams.n_ctx = 256;
|
||||
cparams.n_batch = 256;
|
||||
cparams.n_ubatch = n_ubatch;
|
||||
cparams.kv_unified = false;
|
||||
return llama_init_from_model(model, cparams);
|
||||
};
|
||||
|
||||
llama_context * ctx_roll = make_ctx_multi();
|
||||
llama_context * ctx_ref = make_ctx_multi();
|
||||
if (ctx_roll == nullptr || ctx_ref == nullptr) {
|
||||
fprintf(stderr, "%s : failed to init multi-seq contexts\n", __func__);
|
||||
return false;
|
||||
}
|
||||
|
||||
const auto cleanup = [&]() {
|
||||
llama_free(ctx_roll);
|
||||
llama_free(ctx_ref);
|
||||
};
|
||||
|
||||
if (llama_n_rs_seq(ctx_roll) < n_rollback) {
|
||||
fprintf(stderr, "%s : skipping because n_rs_seq is too small\n", __func__);
|
||||
cleanup();
|
||||
return true;
|
||||
}
|
||||
|
||||
const auto tok = [&](uint32_t seq, llama_pos pos) {
|
||||
return (llama_token) ((7*(uint32_t) pos + 31*seq + 1) % (uint32_t) n_vocab);
|
||||
};
|
||||
|
||||
bool ok = true;
|
||||
|
||||
// both contexts decode the identical [0, p0) prefill; only ctx_roll decodes
|
||||
// the tail, which is then rolled back so its restore is pending at replay
|
||||
for (uint32_t s = 0; s < n_seqs && ok; ++s) {
|
||||
llama_batch batch = llama_batch_init(n_prompt, 0, 1);
|
||||
for (llama_pos pos = 0; pos < (llama_pos) p0; ++pos) {
|
||||
common_batch_add(batch, tok(s, pos), pos, { (llama_seq_id) s }, false);
|
||||
}
|
||||
ok = ok && llama_decode(ctx_roll, batch) == 0;
|
||||
ok = ok && llama_decode(ctx_ref, batch) == 0;
|
||||
|
||||
common_batch_clear(batch);
|
||||
for (llama_pos pos = p0; pos < (llama_pos) n_prompt; ++pos) {
|
||||
common_batch_add(batch, tok(s, pos), pos, { (llama_seq_id) s }, false);
|
||||
}
|
||||
ok = ok && llama_decode(ctx_roll, batch) == 0;
|
||||
llama_batch_free(batch);
|
||||
|
||||
ok = ok && llama_memory_seq_rm(llama_get_memory(ctx_roll), (llama_seq_id) s, p0, -1);
|
||||
|
||||
// a second partial removal while one is pending must be refused
|
||||
ok = ok && !llama_memory_seq_rm(llama_get_memory(ctx_roll), (llama_seq_id) s, p0 - 1, -1);
|
||||
}
|
||||
if (!ok) {
|
||||
fprintf(stderr, "%s : multi-seq prefill/rollback failed\n", __func__);
|
||||
cleanup();
|
||||
return false;
|
||||
}
|
||||
|
||||
llama_batch batch = llama_batch_init(n_seqs*n_replay, 0, 1);
|
||||
for (uint32_t s = 0; s < n_seqs; ++s) {
|
||||
for (uint32_t i = 0; i < n_replay; ++i) {
|
||||
const llama_pos pos = p0 + (llama_pos) i;
|
||||
common_batch_add(batch, tok(s, pos), pos, { (llama_seq_id) s }, true);
|
||||
}
|
||||
}
|
||||
ok = llama_decode(ctx_roll, batch) == 0;
|
||||
ok = ok && llama_decode(ctx_ref, batch) == 0;
|
||||
llama_batch_free(batch);
|
||||
if (!ok) {
|
||||
fprintf(stderr, "%s : multi-seq replay decode failed\n", __func__);
|
||||
cleanup();
|
||||
return false;
|
||||
}
|
||||
|
||||
// identical ubatch shapes from bit-exact states: a correct implementation
|
||||
// matches bitwise, so eps only allows backend scheduling noise
|
||||
constexpr float eps = 1e-7f;
|
||||
|
||||
float diff_max = 0.0f;
|
||||
uint32_t seq_first = 0;
|
||||
int32_t pos_first = -1;
|
||||
for (uint32_t i = 0; i < n_seqs*n_replay; ++i) {
|
||||
const float * l_roll = llama_get_logits_ith(ctx_roll, i);
|
||||
const float * l_ref = llama_get_logits_ith(ctx_ref, i);
|
||||
if (l_roll == nullptr || l_ref == nullptr) {
|
||||
fprintf(stderr, "%s : missing multi-seq logits at index %u\n", __func__, i);
|
||||
cleanup();
|
||||
return false;
|
||||
}
|
||||
for (int t = 0; t < n_vocab; ++t) {
|
||||
const float diff = std::fabs(l_roll[t] - l_ref[t]);
|
||||
if (diff > eps && pos_first < 0) {
|
||||
seq_first = i/n_replay;
|
||||
pos_first = p0 + (int32_t) (i%n_replay);
|
||||
}
|
||||
diff_max = std::max(diff_max, diff);
|
||||
}
|
||||
}
|
||||
|
||||
if (diff_max > eps) {
|
||||
fprintf(stderr, "%s : multi-seq split replay logits mismatch (max diff %g, first at seq %u pos %d)\n",
|
||||
__func__, (double) diff_max, seq_first, pos_first);
|
||||
cleanup();
|
||||
return false;
|
||||
}
|
||||
|
||||
fprintf(stderr, "%s : multi-seq split replay matched (max diff %g)\n", __func__, (double) diff_max);
|
||||
|
||||
// seq-1-only decodes must be independent of seq 0's content: diverge seq 0
|
||||
// in ctx_ref only, then compare identical seq-1-only continuations bitwise
|
||||
constexpr uint32_t n_tail = 4;
|
||||
|
||||
{
|
||||
llama_batch batch_tail = llama_batch_init(n_tail, 0, 1);
|
||||
for (uint32_t i = 0; i < n_tail; ++i) {
|
||||
const llama_pos pos = p0 + (llama_pos) (n_replay + i);
|
||||
common_batch_add(batch_tail, tok(0, pos + 7), pos, { 0 }, false);
|
||||
}
|
||||
ok = llama_decode(ctx_ref, batch_tail) == 0;
|
||||
llama_batch_free(batch_tail);
|
||||
}
|
||||
|
||||
float diff_tail = 0.0f;
|
||||
for (uint32_t i = 0; i < n_tail && ok; ++i) {
|
||||
const llama_pos pos = p0 + (llama_pos) (n_replay + i);
|
||||
llama_batch batch_one = llama_batch_init(1, 0, 1);
|
||||
common_batch_add(batch_one, tok(1, pos), pos, { 1 }, true);
|
||||
ok = llama_decode(ctx_roll, batch_one) == 0;
|
||||
ok = ok && llama_decode(ctx_ref, batch_one) == 0;
|
||||
llama_batch_free(batch_one);
|
||||
if (!ok) {
|
||||
break;
|
||||
}
|
||||
|
||||
const float * l_roll = llama_get_logits_ith(ctx_roll, 0);
|
||||
const float * l_ref = llama_get_logits_ith(ctx_ref, 0);
|
||||
ok = l_roll != nullptr && l_ref != nullptr;
|
||||
for (int t = 0; ok && t < n_vocab; ++t) {
|
||||
diff_tail = std::max(diff_tail, std::fabs(l_roll[t] - l_ref[t]));
|
||||
}
|
||||
}
|
||||
|
||||
if (!ok || diff_tail > eps) {
|
||||
fprintf(stderr, "%s : seq-1-only decode leaked seq 0 state (ok=%d, max diff %g)\n",
|
||||
__func__, ok ? 1 : 0, (double) diff_tail);
|
||||
cleanup();
|
||||
return false;
|
||||
}
|
||||
|
||||
fprintf(stderr, "%s : seq-1-only decode independent of seq 0 (max diff %g)\n", __func__, (double) diff_tail);
|
||||
cleanup();
|
||||
return true;
|
||||
}
|
||||
|
||||
int main(int argc, char ** argv) {
|
||||
std::setlocale(LC_NUMERIC, "C");
|
||||
|
||||
@@ -220,5 +392,10 @@ int main(int argc, char ** argv) {
|
||||
llama_free(ctx_src);
|
||||
llama_free(ctx_dst);
|
||||
llama_free(ctx_dirty);
|
||||
|
||||
if (!test_multi_seq_split_replay(params, model, n_vocab)) {
|
||||
return 1;
|
||||
}
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user