Fix crash with draft-simple (#25720)
* Fix crash with draft-simple * Fix tests for spec decoding
This commit is contained in:
@@ -260,7 +260,10 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
|
|||||||
bool process(const llama_batch & batch) override {
|
bool process(const llama_batch & batch) override {
|
||||||
auto * ctx_dft = params.ctx_dft;
|
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) {
|
if (ret != 0) {
|
||||||
SPC_ERR("failed to decode draft batch, ret = %d\n", ret);
|
SPC_ERR("failed to decode draft batch, ret = %d\n", ret);
|
||||||
|
|||||||
@@ -12,8 +12,9 @@ def create_server():
|
|||||||
server = ServerPreset.stories15m_moe()
|
server = ServerPreset.stories15m_moe()
|
||||||
# set default values
|
# set default values
|
||||||
server.model_draft = download_file(MODEL_DRAFT_FILE_URL)
|
server.model_draft = download_file(MODEL_DRAFT_FILE_URL)
|
||||||
server.draft_min = 4
|
server.spec_type = "draft-simple"
|
||||||
server.draft_max = 8
|
server.spec_draft_n_min = 4
|
||||||
|
server.spec_draft_n_max = 8
|
||||||
server.fa = "off"
|
server.fa = "off"
|
||||||
|
|
||||||
|
|
||||||
@@ -25,6 +26,7 @@ def fixture_create_server():
|
|||||||
def test_with_and_without_draft():
|
def test_with_and_without_draft():
|
||||||
global server
|
global server
|
||||||
server.model_draft = None # disable draft model
|
server.model_draft = None # disable draft model
|
||||||
|
server.spec_type = None
|
||||||
server.start()
|
server.start()
|
||||||
res = server.make_request("POST", "/completion", data={
|
res = server.make_request("POST", "/completion", data={
|
||||||
"prompt": "I believe the meaning of life is",
|
"prompt": "I believe the meaning of life is",
|
||||||
@@ -46,6 +48,7 @@ def test_with_and_without_draft():
|
|||||||
"n_predict": 16,
|
"n_predict": 16,
|
||||||
})
|
})
|
||||||
assert res.status_code == 200
|
assert res.status_code == 200
|
||||||
|
assert res.body["timings"]["draft_n"] > 0
|
||||||
content_draft = res.body["content"]
|
content_draft = res.body["content"]
|
||||||
|
|
||||||
assert content_no_draft == content_draft
|
assert content_no_draft == content_draft
|
||||||
@@ -63,8 +66,8 @@ def test_different_draft_min_draft_max():
|
|||||||
last_content = None
|
last_content = None
|
||||||
for draft_min, draft_max in test_values:
|
for draft_min, draft_max in test_values:
|
||||||
server.stop()
|
server.stop()
|
||||||
server.draft_min = draft_min
|
server.spec_draft_n_min = draft_min
|
||||||
server.draft_max = draft_max
|
server.spec_draft_n_max = draft_max
|
||||||
server.start()
|
server.start()
|
||||||
res = server.make_request("POST", "/completion", data={
|
res = server.make_request("POST", "/completion", data={
|
||||||
"prompt": "I believe the meaning of life is",
|
"prompt": "I believe the meaning of life is",
|
||||||
|
|||||||
@@ -95,6 +95,7 @@ class ServerProcess:
|
|||||||
no_models_autoload: bool | None = None
|
no_models_autoload: bool | None = None
|
||||||
lora_files: List[str] | None = None
|
lora_files: List[str] | None = None
|
||||||
enable_ctx_shift: int | None = False
|
enable_ctx_shift: int | None = False
|
||||||
|
spec_type: str | None = None
|
||||||
spec_draft_n_min: int | None = None
|
spec_draft_n_min: int | None = None
|
||||||
spec_draft_n_max: int | None = None
|
spec_draft_n_max: int | None = None
|
||||||
no_ui: bool | None = None
|
no_ui: bool | None = None
|
||||||
@@ -226,6 +227,8 @@ class ServerProcess:
|
|||||||
server_args.extend(["--lora", lora_file])
|
server_args.extend(["--lora", lora_file])
|
||||||
if self.enable_ctx_shift:
|
if self.enable_ctx_shift:
|
||||||
server_args.append("--context-shift")
|
server_args.append("--context-shift")
|
||||||
|
if self.spec_type:
|
||||||
|
server_args.extend(["--spec-type", self.spec_type])
|
||||||
if self.api_key:
|
if self.api_key:
|
||||||
server_args.extend(["--api-key", self.api_key])
|
server_args.extend(["--api-key", self.api_key])
|
||||||
if self.spec_draft_n_max:
|
if self.spec_draft_n_max:
|
||||||
|
|||||||
Reference in New Issue
Block a user