Skip to content

Commit d2585e5

Browse files
committed
fix: handle Qwen 3.5 hybrid prefix reuse
1 parent e1f8ac0 commit d2585e5

3 files changed

Lines changed: 93 additions & 8 deletions

File tree

llama_cpp/_internals.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -291,10 +291,10 @@ def kv_cache_clear(self):
291291
assert self.memory is not None, "Memory is not initialized"
292292
llama_cpp.llama_memory_clear(self.memory, True)
293293

294-
def kv_cache_seq_rm(self, seq_id: int, p0: int, p1: int):
294+
def kv_cache_seq_rm(self, seq_id: int, p0: int, p1: int) -> bool:
295295
assert self.memory is not None, "Memory is not initialized"
296296
seq_id = seq_id if seq_id >= 0 else 0
297-
llama_cpp.llama_memory_seq_rm(self.memory, seq_id, p0, p1)
297+
return llama_cpp.llama_memory_seq_rm(self.memory, seq_id, p0, p1)
298298

299299
def kv_cache_seq_cp(self, seq_id_src: int, seq_id_dst: int, p0: int, p1: int):
300300
assert self.memory is not None, "Memory is not initialized"

llama_cpp/llama.py

Lines changed: 13 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -891,13 +891,20 @@ def generate(
891891
else:
892892
break
893893
if longest_prefix > 0:
894-
reset = False
895-
tokens = tokens[longest_prefix:]
896-
self.n_tokens = longest_prefix
897-
if self.verbose:
894+
if self._ctx.kv_cache_seq_rm(-1, longest_prefix, -1):
895+
reset = False
896+
tokens = tokens[longest_prefix:]
897+
self.n_tokens = longest_prefix
898+
if self.verbose:
899+
print(
900+
f"Llama.generate: {longest_prefix} prefix-match hit, "
901+
f"remaining {len(tokens)} prompt tokens to eval",
902+
file=sys.stderr,
903+
)
904+
elif self.verbose:
898905
print(
899-
f"Llama.generate: {longest_prefix} prefix-match hit, "
900-
f"remaining {len(tokens)} prompt tokens to eval",
906+
f"Llama.generate: {longest_prefix} prefix-match found "
907+
f"but partial kv removal not supported, re-evaluating full prompt",
901908
file=sys.stderr,
902909
)
903910

tests/test_llama.py

Lines changed: 78 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
import ctypes
22
import multiprocessing
3+
from types import SimpleNamespace
4+
from unittest.mock import Mock
35

46
import numpy as np
57
from scipy.special import log_softmax
@@ -15,6 +17,10 @@
1517
MODEL = "./vendor/llama.cpp/models/ggml-vocab-llama-spm.gguf"
1618

1719

20+
class EvalCalled(Exception):
21+
pass
22+
23+
1824
def test_llama_cpp_version():
1925
assert llama_cpp.__version__
2026

@@ -232,3 +238,75 @@ def test_real_llama_embeddings(llama_cpp_model_path):
232238
)
233239
# Smoke test for now
234240
model.embed("Hello World")
241+
242+
243+
def test_kv_cache_seq_rm_returns_bool(monkeypatch):
244+
context = internals.LlamaContext.__new__(internals.LlamaContext)
245+
context.memory = object()
246+
calls = []
247+
248+
def fake_llama_memory_seq_rm(memory, seq_id, p0, p1):
249+
calls.append((memory, seq_id, p0, p1))
250+
return True
251+
252+
monkeypatch.setattr(llama_cpp, "llama_memory_seq_rm", fake_llama_memory_seq_rm)
253+
254+
assert context.kv_cache_seq_rm(-1, 4, -1) is True
255+
assert calls == [(context.memory, 0, 4, -1)]
256+
257+
258+
def make_test_llama(kv_cache_seq_rm_return):
259+
llama = llama_cpp.Llama.__new__(llama_cpp.Llama)
260+
llama.n_tokens = 3
261+
llama.n_batch = 8
262+
llama._n_ctx = 32
263+
llama._n_vocab = 8
264+
llama._logits_all = False
265+
llama._seed = 1337
266+
llama.last_n_tokens_size = 64
267+
llama.verbose = False
268+
llama.input_ids = np.array([1, 2, 3, 0, 0, 0], dtype=np.intc)
269+
llama.scores = np.zeros((6, 8), dtype=np.single)
270+
llama._ctx = SimpleNamespace(
271+
kv_cache_seq_rm=Mock(return_value=kv_cache_seq_rm_return)
272+
)
273+
llama._sampler = None
274+
llama.eval_tokens_seen = None
275+
llama.reset_calls = 0
276+
277+
def reset():
278+
llama.reset_calls += 1
279+
llama.n_tokens = 0
280+
281+
def eval_tokens(tokens):
282+
llama.eval_tokens_seen = list(tokens)
283+
raise EvalCalled
284+
285+
llama.reset = reset
286+
llama.eval = eval_tokens
287+
llama._init_sampler = lambda **kwargs: object()
288+
return llama
289+
290+
291+
def test_generate_reuses_prefix_when_partial_removal_supported():
292+
llama = make_test_llama(True)
293+
294+
with pytest.raises(EvalCalled):
295+
next(llama.generate([1, 2, 3, 4]))
296+
297+
llama._ctx.kv_cache_seq_rm.assert_called_once_with(-1, 3, -1)
298+
assert llama.reset_calls == 0
299+
assert llama.n_tokens == 3
300+
assert llama.eval_tokens_seen == [4]
301+
302+
303+
def test_generate_falls_back_to_reset_when_partial_removal_rejected():
304+
llama = make_test_llama(False)
305+
306+
with pytest.raises(EvalCalled):
307+
next(llama.generate([1, 2, 3, 4]))
308+
309+
llama._ctx.kv_cache_seq_rm.assert_called_once_with(-1, 3, -1)
310+
assert llama.reset_calls == 1
311+
assert llama.n_tokens == 0
312+
assert llama.eval_tokens_seen == [1, 2, 3, 4]

0 commit comments

Comments
 (0)