spec: Add benchmark-only synthetic speculative acceptance options (#27711)

* Add benchmark-only synthetic speculative acceptance to llama-server and llama-cli

* Address review comments

* Address review comments

* Add some comments in the code
This commit is contained in:
Gaurav Garg
2026-08-27 13:53:42 +03:00
committed by GitHub
parent deae5ee133
commit 2bb9bddafa
11 changed files with 427 additions and 17 deletions
@@ -52,6 +52,18 @@ def test_with_and_without_draft():
assert tokens_no_draft == tokens_draft
server.stop()
create_server()
assert server.spec_draft_n_max is not None
server.spec_synth_rates = [0.0] * server.spec_draft_n_max
server.start()
res = server.make_request("POST", "/completion", data=request)
assert res.status_code == 200
assert res.body["timings"]["draft_n"] > 0
assert res.body["timings"]["draft_n_accepted"] == 0
assert res.body["tokens"] == tokens_no_draft
def test_different_draft_min_draft_max():
global server
@@ -80,6 +92,66 @@ def test_different_draft_min_draft_max():
last_content = res.body["content"]
def test_synth_is_deterministic():
global server
assert server.spec_draft_n_max is not None
server.spec_synth_rates = [0.75 ** (i + 1) for i in range(server.spec_draft_n_max)]
server.start()
request = {
"prompt": "I believe the meaning of life is",
"temperature": 0.2,
"top_k": 5,
"seed": 4242,
"n_predict": 32,
}
responses = [server.make_request("POST", "/completion", data=request) for _ in range(2)]
for res in responses:
assert res.status_code == 200
assert res.body["timings"]["draft_n"] > 0
assert responses[0].body["timings"]["draft_n"] == responses[1].body["timings"]["draft_n"]
assert responses[0].body["timings"]["draft_n_accepted"] == responses[1].body["timings"]["draft_n_accepted"]
def test_synth_ignores_target_tokens():
global server
assert server.spec_draft_n_max is not None
server.spec_synth_rates = [1.0] * server.spec_draft_n_max
server.start()
res = server.make_request("POST", "/completion", data={
"prompt": "I believe the meaning of life is",
"temperature": 0.0,
"seed": 4242,
"n_predict": 32,
})
assert res.status_code == 200
assert res.body["timings"]["draft_n"] > 0
assert res.body["timings"]["draft_n_accepted"] == res.body["timings"]["draft_n"]
res = server.make_request("POST", "/completion", data={
"prompt": "I believe the meaning of life is",
"temperature": 0.0,
"seed": 4242,
"n_predict": 6,
"grammar": 'root ::= "a"{5,5}',
})
assert res.status_code == 200, res.body
res = server.make_request("POST", "/completion", data={
"prompt": "Respond with only: OK",
"temperature": 0.0,
"seed": 4242,
"n_predict": 64,
"ignore_eos": True,
})
assert res.status_code == 200, res.body
assert res.body["tokens_predicted"] == 64
assert res.body["stop_type"] == "limit"
def test_slot_ctx_not_exceeded():
global server
server.n_ctx = 256
+7
View File
@@ -99,6 +99,8 @@ class ServerProcess:
spec_type: str | None = None
spec_draft_n_min: int | None = None
spec_draft_n_max: int | None = None
spec_synth_len: float | None = None
spec_synth_rates: List[float] | None = None
no_ui: bool | None = None
jinja: bool | None = None
reasoning_format: Literal['deepseek', 'none', 'nothink'] | None = None
@@ -245,6 +247,11 @@ class ServerProcess:
server_args.extend(["--spec-draft-n-max", self.spec_draft_n_max])
if self.spec_draft_n_min:
server_args.extend(["--spec-draft-n-min", self.spec_draft_n_min])
if self.spec_synth_len is not None:
server_args.extend(["--spec-synth-len", self.spec_synth_len])
if self.spec_synth_rates is not None:
rates = ",".join(str(rate) for rate in self.spec_synth_rates)
server_args.extend(["--spec-synth-rates", rates])
if self.no_ui:
server_args.append("--no-ui")
if self.no_models_autoload: