Fix crash with draft-simple (#25720)

* Fix crash with draft-simple

* Fix tests for spec decoding
This commit is contained in:
Gaurav Garg
2026-07-15 19:51:34 +05:30
committed by GitHub
parent a582222290
commit 956973c764
3 changed files with 14 additions and 5 deletions
+7 -4
View 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",
+3
View 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: