Commit 4f31296a9 for llama.cpp
commit 4f31296a907d73f3409f720429bcf31111615d77
Author: Sihan Yu <116848194+sihanyu03@users.noreply.github.com>
Date: Wed Sep 30 20:39:53 2026 +0300
test-llama-archs : toggle causal_attn to catch graph shape changes (#29724)
After the device decode, flip causal_attn off, decode n_ubatch/2 then
n_ubatch tokens. Both have the same node count, so a shape that depends
on the flag makes the second reallocate at an unchanged graph size,
which aborts under GGML_SCHED_NO_REALLOC. Skipped for the encode archs.
diff --git a/tests/test-llama-archs.cpp b/tests/test-llama-archs.cpp
index 015e3414b..e3731e6f8 100644
--- a/tests/test-llama-archs.cpp
+++ b/tests/test-llama-archs.cpp
@@ -535,6 +535,51 @@ static std::vector<float> get_logits(
return ret;
}
+static bool check_causal_attn_toggle(
+ llama_model * model, llama_context * lctx, const std::vector<llama_token> & tokens) {
+ const uint32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(model));
+ const uint32_t n_past = tokens.size();
+ const uint32_t n_ubatch = llama_n_ubatch(lctx);
+
+ GGML_ASSERT(n_past + n_ubatch/2 + n_ubatch <= llama_n_ctx(lctx));
+
+ llama_set_causal_attn(lctx, false);
+
+ common_batch batch(lctx);
+
+ bool ok = true;
+ uint32_t pos = n_past;
+ for (const uint32_t n_tokens : { n_ubatch/2, n_ubatch }) {
+ batch.clear();
+ for (uint32_t i = 0; i < n_tokens; i++) {
+ batch.add(tokens[i], pos++, 0, true);
+ }
+
+ const int32_t rc = llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+ if (rc != 0) {
+ LOG_ERR("%s: n_tokens=%u: llama_process returned %d\n", __func__, n_tokens, rc);
+ ok = false;
+ break;
+ }
+
+ const float * logits = llama_get_logits_ith(lctx, n_tokens - 1);
+ if (logits == nullptr) {
+ LOG_ERR("%s: n_tokens=%u: no logits\n", __func__, n_tokens);
+ ok = false;
+ break;
+ }
+ for (uint32_t j = 0; j < n_vocab; j++) {
+ if (std::isnan(logits[j])) {
+ LOG_ERR("%s: n_tokens=%u: nan logit\n", __func__, n_tokens);
+ ok = false;
+ break;
+ }
+ }
+ }
+
+ return ok;
+}
+
static bool moe_mandatory(const llm_arch arch) {
switch (arch) {
case LLM_ARCH_LLAMA4:
@@ -862,6 +907,12 @@ static int test_backends(const std::string & arch_filter, const size_t seed, con
test_ok = false;
status_nmse = "\033[1;31mFAIL\033[0m";
}
+ if (!encode && !check_causal_attn_toggle(model_and_ctx_dev.first.get(), model_and_ctx_dev.second.get(), tokens)) {
+ if (test_ok) {
+ status_nmse = "\033[1;31mFAIL\033[0m (toggle)";
+ }
+ test_ok = false;
+ }
}
FILE * file = tmpfile(); // Can be null on Windows without administrator privileges.