]> git.djapps.eu Git - pkg/ggml/sources/llama.cpp/commitdiff
Fix crash with draft-simple (#25720)
authorGaurav Garg <redacted>
Wed, 15 Jul 2026 14:21:34 +0000 (19:51 +0530)
committerGitHub <redacted>
Wed, 15 Jul 2026 14:21:34 +0000 (19:51 +0530)
* Fix crash with draft-simple

* Fix tests for spec decoding

common/speculative.cpp
tools/server/tests/unit/test_speculative.py
tools/server/tests/utils.py

index 580728a2001eb1171ed2ece5041c80b575d0a70f..3cb08767bd46bc099a7aab159c632b8cfdcdba8b 100644 (file)
@@ -260,7 +260,10 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
     bool process(const llama_batch & batch) override {
         auto * ctx_dft = params.ctx_dft;
 
-        const int ret = llama_decode(ctx_dft, batch);
+        llama_batch batch_dft = batch;
+        batch_dft.logits = nullptr;
+
+        const int ret = llama_decode(ctx_dft, batch_dft);
 
         if (ret != 0) {
             SPC_ERR("failed to decode draft batch, ret = %d\n", ret);
index 84cd77e6f2ed59aef4e29dd739b5df916f3cf3f9..c6568479ca4ab9fd8b31c03cbc5aa98a8fd6ced2 100644 (file)
@@ -12,8 +12,9 @@ def create_server():
     server = ServerPreset.stories15m_moe()
     # set default values
     server.model_draft = download_file(MODEL_DRAFT_FILE_URL)
-    server.draft_min = 4
-    server.draft_max = 8
+    server.spec_type = "draft-simple"
+    server.spec_draft_n_min = 4
+    server.spec_draft_n_max = 8
     server.fa = "off"
 
 
@@ -25,6 +26,7 @@ def fixture_create_server():
 def test_with_and_without_draft():
     global server
     server.model_draft = None  # disable draft model
+    server.spec_type = None
     server.start()
     res = server.make_request("POST", "/completion", data={
         "prompt": "I believe the meaning of life is",
@@ -46,6 +48,7 @@ def test_with_and_without_draft():
         "n_predict": 16,
     })
     assert res.status_code == 200
+    assert res.body["timings"]["draft_n"] > 0
     content_draft = res.body["content"]
 
     assert content_no_draft == content_draft
@@ -63,8 +66,8 @@ def test_different_draft_min_draft_max():
     last_content = None
     for draft_min, draft_max in test_values:
         server.stop()
-        server.draft_min = draft_min
-        server.draft_max = draft_max
+        server.spec_draft_n_min = draft_min
+        server.spec_draft_n_max = draft_max
         server.start()
         res = server.make_request("POST", "/completion", data={
             "prompt": "I believe the meaning of life is",
index f4f0e61e6106ce4e7355b3088d916758da2c0cdd..5d5c873ac4cceed3acab39ce01639a0bb5844c5e 100644 (file)
@@ -95,6 +95,7 @@ class ServerProcess:
     no_models_autoload: bool | None = None
     lora_files: List[str] | None = None
     enable_ctx_shift: int | None = False
+    spec_type: str | None = None
     spec_draft_n_min: int | None = None
     spec_draft_n_max: int | None = None
     no_ui: bool | None = None
@@ -226,6 +227,8 @@ class ServerProcess:
                 server_args.extend(["--lora", lora_file])
         if self.enable_ctx_shift:
             server_args.append("--context-shift")
+        if self.spec_type:
+            server_args.extend(["--spec-type", self.spec_type])
         if self.api_key:
             server_args.extend(["--api-key", self.api_key])
         if self.spec_draft_n_max: