Felipe97/llama-cpp-compiled
01.2k
1import pytest2import requests3import time4import random5 6from openai import OpenAI7from utils import *8 9server = ServerPreset.tinyllama2()10 11JSON_MULTIMODAL_KEY = "multimodal_data"12JSON_PROMPT_STRING_KEY = "prompt_string"13 14@pytest.fixture(autouse=True)15def create_server():16 global server17 server = ServerPreset.tinyllama2()18 19@pytest.mark.parametrize("prompt,n_predict,re_content,n_prompt,n_predicted,truncated,return_tokens", [20 ("I believe the meaning of life is", 8, "(going|bed)+", 18, 8, False, False),21 ("Write a joke about AI from a very long prompt which will not be truncated", 64, "(princesses|everyone|kids|Anna|forest)+", 46, 64, False, True),22])23def test_completion(prompt: str, n_predict: int, re_content: str, n_prompt: int, n_predicted: int, truncated: bool, return_tokens: bool):24 global server25 server.start()26 res = server.make_request("POST", "/completion", data={27 "n_predict": n_predict,28 "prompt": prompt,29 "return_tokens": return_tokens,30 })31 assert res.status_code == 20032 assert res.body["timings"]["prompt_n"] == n_prompt33 assert res.body["timings"]["predicted_n"] == n_predicted34 assert res.body["truncated"] == truncated35 assert type(res.body["has_new_line"]) == bool36 assert match_regex(re_content, res.body["content"])37 if return_tokens:38 assert len(res.body["tokens"]) > 039 assert all(type(tok) == int for tok in res.body["tokens"])40 else:41 assert res.body["tokens"] == []42 43 44@pytest.mark.parametrize("prompt,n_predict,re_content,n_prompt,n_predicted,truncated", [45 ("I believe the meaning of life is", 8, "(going|bed)+", 18, 8, False),46 ("Write a joke about AI from a very long prompt which will not be truncated", 64, "(princesses|everyone|kids|Anna|forest)+", 46, 64, False),47])48def test_completion_stream(prompt: str, n_predict: int, re_content: str, n_prompt: int, n_predicted: int, truncated: bool):49 global server50 server.start()51 res = server.make_stream_request("POST", "/completion", data={52 "n_predict": n_predict,53 "prompt": prompt,54 "stream": True,55 })56 content = ""57 for data in res:58 assert "stop" in data and type(data["stop"]) == bool59 if data["stop"]:60 assert data["timings"]["prompt_n"] == n_prompt61 assert data["timings"]["predicted_n"] == n_predicted62 assert data["truncated"] == truncated63 assert data["stop_type"] == "limit"64 assert type(data["has_new_line"]) == bool65 assert "generation_settings" in data66 assert server.n_predict is not None67 assert data["generation_settings"]["n_predict"] == min(n_predict, server.n_predict)68 assert data["generation_settings"]["seed"] == server.seed69 assert "adaptive_target" in data["generation_settings"]70 assert "adaptive_decay" in data["generation_settings"]71 assert match_regex(re_content, content)72 else:73 assert len(data["tokens"]) > 074 assert all(type(tok) == int for tok in data["tokens"])75 content += data["content"]76 77 78def test_completion_stream_vs_non_stream():79 global server80 server.start()81 res_stream = server.make_stream_request("POST", "/completion", data={82 "n_predict": 8,83 "prompt": "I believe the meaning of life is",84 "stream": True,85 })86 res_non_stream = server.make_request("POST", "/completion", data={87 "n_predict": 8,88 "prompt": "I believe the meaning of life is",89 })90 content_stream = ""91 for data in res_stream:92 content_stream += data["content"]93 assert content_stream == res_non_stream.body["content"]94 95 96def test_completion_with_openai_library():97 global server98 server.start()99 client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")100 res = client.completions.create(101 model="davinci-002",102 prompt="I believe the meaning of life is",103 max_tokens=8,104 )105 assert res.system_fingerprint is not None and res.system_fingerprint.startswith("b")106 assert res.choices[0].finish_reason == "length"107 assert res.choices[0].text is not None108 assert match_regex("(going|bed)+", res.choices[0].text)109 110 111def test_completion_stream_with_openai_library():112 global server113 server.start()114 client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")115 res = client.completions.create(116 model="davinci-002",117 prompt="I believe the meaning of life is",118 max_tokens=8,119 stream=True,120 )121 output_text = ''122 for data in res:123 choice = data.choices[0]124 if choice.finish_reason is None:125 assert choice.text is not None126 output_text += choice.text127 assert match_regex("(going|bed)+", output_text)128 129 130# Test case from https://github.com/ggml-org/llama.cpp/issues/13780131@pytest.mark.slow132def test_completion_stream_with_openai_library_stops():133 global server134 server.model_hf_repo = "bartowski/Phi-3.5-mini-instruct-GGUF:Q4_K_M"135 server.model_hf_file = None136 server.start()137 client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")138 res = client.completions.create(139 model="davinci-002",140 prompt="System: You are helpful assistant.\nAssistant:\nHey! How could I help?\nUser:\nTell me a joke.\nAssistant:\n",141 stop=["User:\n", "Assistant:\n"],142 max_tokens=200,143 stream=True,144 )145 output_text = ''146 for data in res:147 choice = data.choices[0]148 if choice.finish_reason is None:149 assert choice.text is not None150 output_text += choice.text151 assert match_regex("Sure, here's one for[\\s\\S]*", output_text), f'Unexpected output: {output_text}'152 153 154@pytest.mark.parametrize("n_slots", [1, 2])155def test_consistent_result_same_seed(n_slots: int):156 global server157 server.n_slots = n_slots158 server.start()159 last_res = None160 for _ in range(4):161 res = server.make_request("POST", "/completion", data={162 "prompt": "I believe the meaning of life is",163 "seed": 42,164 "temperature": 0.0,165 "cache_prompt": False, # TODO: remove this once test_cache_vs_nocache_prompt is fixed166 })167 if last_res is not None:168 assert res.body["content"] == last_res.body["content"]169 last_res = res170 171 172@pytest.mark.parametrize("n_slots", [1, 2])173def test_different_result_different_seed(n_slots: int):174 global server175 server.n_slots = n_slots176 server.start()177 last_res = None178 for seed in range(4):179 res = server.make_request("POST", "/completion", data={180 "prompt": "I believe the meaning of life is",181 "seed": seed,182 "temperature": 1.0,183 "cache_prompt": False, # TODO: remove this once test_cache_vs_nocache_prompt is fixed184 })185 if last_res is not None:186 assert res.body["content"] != last_res.body["content"]187 last_res = res188 189# TODO figure why it don't work with temperature = 1190# @pytest.mark.parametrize("temperature", [0.0, 1.0])191@pytest.mark.parametrize("n_batch", [16, 32])192@pytest.mark.parametrize("temperature", [0.0])193def test_consistent_result_different_batch_size(n_batch: int, temperature: float):194 global server195 server.n_batch = n_batch196 server.start()197 last_res = None198 for _ in range(4):199 res = server.make_request("POST", "/completion", data={200 "prompt": "I believe the meaning of life is",201 "seed": 42,202 "temperature": temperature,203 "cache_prompt": False, # TODO: remove this once test_cache_vs_nocache_prompt is fixed204 })205 if last_res is not None:206 assert res.body["content"] == last_res.body["content"]207 last_res = res208 209 210@pytest.mark.skip(reason="This test fails on linux, need to be fixed")211def test_cache_vs_nocache_prompt():212 global server213 server.start()214 res_cache = server.make_request("POST", "/completion", data={215 "prompt": "I believe the meaning of life is",216 "seed": 42,217 "temperature": 1.0,218 "cache_prompt": True,219 })220 res_no_cache = server.make_request("POST", "/completion", data={221 "prompt": "I believe the meaning of life is",222 "seed": 42,223 "temperature": 1.0,224 "cache_prompt": False,225 })226 assert res_cache.body["content"] == res_no_cache.body["content"]227 228 229def test_nocache_long_input_prompt():230 global server231 server.start()232 res = server.make_request("POST", "/completion", data={233 "prompt": "I believe the meaning of life is"*32,234 "seed": 42,235 "temperature": 1.0,236 "cache_prompt": False,237 })238 assert res.status_code == 400239 240def test_json_prompt_no_mtmd():241 global server242 server.start()243 res = server.make_request("POST", "/completion", data={244 "prompt": { JSON_PROMPT_STRING_KEY: "I believe the meaning of life is" },245 "seed": 42,246 "temperature": 1.0,247 "cache_prompt": False,248 })249 assert res.status_code == 200250 251def test_json_prompt_mtm_error_when_not_supported():252 global server253 server.start()254 res = server.make_request("POST", "/completion", data={255 "prompt": { JSON_PROMPT_STRING_KEY: "I believe the meaning of life is <__media__>", JSON_MULTIMODAL_KEY: "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk+A8AAQUBAScY42YAAAAASUVORK5CYII=" },256 "seed": 42,257 "temperature": 1.0,258 "cache_prompt": False,259 })260 # MTMD is disabled on this model, so this should fail.261 assert res.status_code != 200262 263def test_completion_with_tokens_input():264 global server265 server.temperature = 0.0266 server.start()267 prompt_str = "I believe the meaning of life is"268 res = server.make_request("POST", "/tokenize", data={269 "content": prompt_str,270 "add_special": True,271 })272 assert res.status_code == 200273 tokens = res.body["tokens"]274 275 # single completion276 res = server.make_request("POST", "/completion", data={277 "prompt": tokens,278 })279 assert res.status_code == 200280 assert type(res.body["content"]) == str281 282 # batch completion283 res = server.make_request("POST", "/completion", data={284 "prompt": [tokens, tokens],285 })286 assert res.status_code == 200287 assert type(res.body) == list288 assert len(res.body) == 2289 assert res.body[0]["content"] == res.body[1]["content"]290 291 # mixed string and tokens292 res = server.make_request("POST", "/completion", data={293 "prompt": [tokens, prompt_str],294 })295 assert res.status_code == 200296 assert type(res.body) == list297 assert len(res.body) == 2298 assert res.body[0]["content"] == res.body[1]["content"]299 300 # mixed JSON and tokens301 res = server.make_request("POST", "/completion", data={302 "prompt": [303 tokens,304 {305 JSON_PROMPT_STRING_KEY: "I believe the meaning of life is",306 },307 ],308 })309 assert res.status_code == 200310 assert type(res.body) == list311 assert len(res.body) == 2312 assert res.body[0]["content"] == res.body[1]["content"]313 314 # mixed string and tokens in one sequence315 res = server.make_request("POST", "/completion", data={316 "prompt": [1, 2, 3, 4, 5, 6, prompt_str, 7, 8, 9, 10, prompt_str],317 })318 assert res.status_code == 200319 assert type(res.body["content"]) == str320 321 322@pytest.mark.parametrize("n_slots,n_requests", [323 (1, 3),324 (2, 2),325 (2, 4),326 (4, 2), # some slots must be idle327 (4, 6),328])329def test_completion_parallel_slots(n_slots: int, n_requests: int):330 global server331 server.n_slots = n_slots332 server.temperature = 0.0333 server.start()334 335 PROMPTS = [336 ("Write a very long book.", "(very|special|big)+"),337 ("Write another a poem.", "(small|house)+"),338 ("What is LLM?", "(Dad|said)+"),339 ("The sky is blue and I love it.", "(climb|leaf)+"),340 ("Write another very long music lyrics.", "(friends|step|sky)+"),341 ("Write a very long joke.", "(cat|Whiskers)+"),342 ]343 def check_slots_status():344 should_all_slots_busy = n_requests >= n_slots345 time.sleep(0.1)346 res = server.make_request("GET", "/slots")347 n_busy = sum([1 for slot in res.body if slot["is_processing"]])348 if should_all_slots_busy:349 assert n_busy == n_slots350 else:351 assert n_busy <= n_slots352 353 tasks = []354 for i in range(n_requests):355 prompt, re_content = PROMPTS[i % len(PROMPTS)]356 tasks.append((server.make_request, ("POST", "/completion", {357 "prompt": prompt,358 "seed": 42,359 "temperature": 1.0,360 })))361 tasks.append((check_slots_status, ()))362 results = parallel_function_calls(tasks)363 364 # check results365 for i in range(n_requests):366 prompt, re_content = PROMPTS[i % len(PROMPTS)]367 res = results[i]368 assert res.status_code == 200369 assert type(res.body["content"]) == str370 assert len(res.body["content"]) > 10371 # FIXME: the result is not deterministic when using other slot than slot 0372 # assert match_regex(re_content, res.body["content"])373 374 375@pytest.mark.parametrize(376 "n_ctx,n_slots,n_predict_vals,expected_success",377 [378 (256, 4, [80, 40, 80, 80], [True, True, True, True]),379 (256, 4, [70, 70, 70, 70], [False, False, False, False]),380 (256, 4, [90, 90, 40, 90], [False, False, True, False]),381 (256, 4, [90, 90, 40, 75], [True, True, True, True]),382 ],383)384def test_completion_unified(n_ctx, n_slots, n_predict_vals, expected_success):385 global server386 server.n_slots = n_slots387 server.kv_unified = True388 server.n_ctx = n_ctx389 server.start()390 prompt = "A"391 tasks = []392 for n_predict in n_predict_vals:393 tasks.append((server.make_request, ("POST", "/completion", {"prompt": prompt, "n_predict": n_predict})))394 results = parallel_function_calls(tasks)395 for res, n_predict, expect_ok in zip(results, n_predict_vals, expected_success):396 if expect_ok:397 # the pool is aborted as a whole, so a request that fits on its own398 # is still dropped when the slots overlap, and it says so explicitly399 assert res.status_code == 200 or (400 res.status_code == 500401 and "context size has been exceeded" in res.body["error"]["message"].lower()402 )403 404 # note: https://github.com/ggml-org/llama.cpp/pull/18700#issuecomment-3728695581405 if res.status_code == 200:406 assert "content" in res.body407 if "timings" in res.body:408 assert res.body["timings"]["predicted_n"] == n_predict409 410 411@pytest.mark.parametrize(412 "prompt,n_predict,response_fields",413 [414 ("I believe the meaning of life is", 8, []),415 ("I believe the meaning of life is", 32, ["content", "generation_settings/n_predict", "prompt"]),416 ],417)418def test_completion_response_fields(419 prompt: str, n_predict: int, response_fields: list[str]420):421 global server422 server.start()423 res = server.make_request(424 "POST",425 "/completion",426 data={427 "n_predict": n_predict,428 "prompt": prompt,429 "response_fields": response_fields,430 },431 )432 assert res.status_code == 200433 assert "content" in res.body434 assert len(res.body["content"])435 if len(response_fields):436 assert res.body["generation_settings/n_predict"] == n_predict437 assert res.body["prompt"] == "<s> " + prompt438 assert isinstance(res.body["content"], str)439 assert len(res.body) == len(response_fields)440 else:441 assert len(res.body)442 assert "generation_settings" in res.body443 444 445def test_n_probs():446 global server447 server.start()448 res = server.make_request("POST", "/completion", data={449 "prompt": "I believe the meaning of life is",450 "n_probs": 10,451 "temperature": 0.0,452 "n_predict": 5,453 })454 assert res.status_code == 200455 assert "completion_probabilities" in res.body456 assert len(res.body["completion_probabilities"]) == 5457 for tok in res.body["completion_probabilities"]:458 assert "id" in tok and tok["id"] > 0459 assert "token" in tok and type(tok["token"]) == str460 assert "logprob" in tok and tok["logprob"] <= 0.0461 assert "bytes" in tok and type(tok["bytes"]) == list462 assert len(tok["top_logprobs"]) == 10463 for prob in tok["top_logprobs"]:464 assert "id" in prob and prob["id"] > 0465 assert "token" in prob and type(prob["token"]) == str466 assert "logprob" in prob and prob["logprob"] <= 0.0467 assert "bytes" in prob and type(prob["bytes"]) == list468 469 470def test_n_probs_stream():471 global server472 server.start()473 res = server.make_stream_request("POST", "/completion", data={474 "prompt": "I believe the meaning of life is",475 "n_probs": 10,476 "temperature": 0.0,477 "n_predict": 5,478 "stream": True,479 })480 for data in res:481 if data["stop"] == False:482 assert "completion_probabilities" in data483 assert len(data["completion_probabilities"]) == 1484 for tok in data["completion_probabilities"]:485 assert "id" in tok and tok["id"] > 0486 assert "token" in tok and type(tok["token"]) == str487 assert "logprob" in tok and tok["logprob"] <= 0.0488 assert "bytes" in tok and type(tok["bytes"]) == list489 assert len(tok["top_logprobs"]) == 10490 for prob in tok["top_logprobs"]:491 assert "id" in prob and prob["id"] > 0492 assert "token" in prob and type(prob["token"]) == str493 assert "logprob" in prob and prob["logprob"] <= 0.0494 assert "bytes" in prob and type(prob["bytes"]) == list495 496 497def test_n_probs_post_sampling():498 global server499 server.start()500 res = server.make_request("POST", "/completion", data={501 "prompt": "Today was the day. Today I would finally become a",502 "n_probs": 10,503 "temperature": 1.0,504 "n_predict": 5,505 "post_sampling_probs": True,506 })507 assert res.status_code == 200508 assert "completion_probabilities" in res.body509 assert len(res.body["completion_probabilities"]) == 5510 for (i, tok) in enumerate(res.body["completion_probabilities"]):511 assert "id" in tok and tok["id"] > 0512 assert "token" in tok and type(tok["token"]) == str513 assert "prob" in tok and 0.0 < tok["prob"] <= 1.0514 assert "bytes" in tok and type(tok["bytes"]) == list515 assert "top_probs" in tok and type(tok["top_probs"]) == list516 517 for prob in tok["top_probs"]:518 assert "id" in prob and prob["id"] > 0519 assert "token" in prob and type(prob["token"]) == str520 # 0.0 probability tokens should never be returned by the server521 assert "prob" in prob and 0.0 < prob["prob"] <= 1.0522 assert "bytes" in prob and type(prob["bytes"]) == list523 524 if i == 0:525 # The prompt is vague enough that we should get at least 10 possibilities526 # for the first token.527 assert len(tok["top_probs"]) == 10528 529 if len(tok["top_probs"]) < 10:530 # Getting less than the requested number of probabilities should only happen531 # if the ones we did get already sum to 1.0.532 assert sum(p["prob"] for p in tok["top_probs"]) == pytest.approx(1.0)533 534def test_n_probs_post_backend_sampling():535 """Verify that the same probabilities are returned with and without backend sampling."""536 global server537 server.backend_sampling = True538 server.start()539 540 def make_request(backend_sampling):541 n_predict = 20542 543 res = server.make_request("POST", "/completion", data={544 "prompt": "The countries of Europe, in random order, are:",545 "n_probs": 10,546 "n_predict": n_predict,547 "post_sampling_probs": True,548 "seed": 4242,549 "backend_sampling": backend_sampling,550 })551 assert res.status_code == 200552 553 total_probs = 0554 completions = res.body["completion_probabilities"]555 assert len(completions) == n_predict556 for tok in completions:557 # Handling of 0.0 probabilities differs between samplers and backend sampling. Filter them to normalize the558 # data.559 tok["top_probs"] = [x for x in tok["top_probs"] if x["prob"] > 0.0]560 total_probs += len(tok["top_probs"])561 # Verify that we got at least two top probs on average, to ensure the effectiveness of the test.562 assert total_probs >= 2 * n_predict563 return completions564 565 def verify_token(a, b):566 assert a["id"] == b["id"]567 assert a["token"] == b["token"]568 assert a["bytes"] == b["bytes"]569 assert a["prob"] == pytest.approx(b["prob"], abs=0.01)570 571 for (a, b) in zip(make_request(True), make_request(False)):572 verify_token(a, b)573 assert len(a["top_probs"]) == len(b["top_probs"])574 575 for (aa, bb) in zip(a["top_probs"], b["top_probs"]):576 verify_token(aa, bb)577 578@pytest.mark.parametrize("tokenize,openai_style", [(False, False), (False, True), (True, False), (True, True)])579def test_logit_bias(tokenize, openai_style):580 global server581 server.start()582 583 exclude = ["i", "I", "the", "The", "to", "a", "an", "be", "is", "was", "but", "But", "and", "And", "so", "So", "you", "You", "he", "He", "she", "She", "we", "We", "they", "They", "it", "It", "his", "His", "her", "Her", "book", "Book"]584 585 logit_bias = []586 if tokenize:587 res = server.make_request("POST", "/tokenize", data={588 "content": " " + " ".join(exclude) + " ",589 })590 assert res.status_code == 200591 tokens = res.body["tokens"]592 logit_bias = [[tok, -100] for tok in tokens]593 594 else:595 logit_bias = [[" " + tok + " ", -100] for tok in exclude]596 597 if openai_style:598 logit_bias = {el[0]: -100 for el in logit_bias}599 600 res = server.make_request("POST", "/completion", data={601 "n_predict": 64,602 "prompt": "What is the best book",603 "logit_bias": logit_bias,604 "temperature": 0.0605 })606 assert res.status_code == 200607 output_text = res.body["content"]608 assert all(output_text.find(" " + tok + " ") == -1 for tok in exclude)609 610 611def test_cancel_request():612 global server613 server.n_ctx = 4096614 server.n_predict = -1615 server.n_slots = 1616 server.server_slots = True617 server.start()618 # send a request that will take a long time, but cancel it before it finishes619 try:620 server.make_request("POST", "/completion", data={621 "prompt": "I believe the meaning of life is",622 }, timeout=0.1)623 except requests.exceptions.ReadTimeout:624 pass # expected625 # make sure the slot is free626 time.sleep(2)627 res = server.make_request("GET", "/slots")628 assert res.body[0]["is_processing"] == False629 630 631# this test exercises the host-memory prompt cache632# ref: https://github.com/ggml-org/llama.cpp/pull/16391633# ref: https://github.com/ggml-org/llama.cpp/pull/17078634def test_completion_prompt_cache():635 global server636 server.n_slots = 2637 server.kv_unified = True638 server.start()639 640 for _ in range(16):641 # generate alternating random prompts with variable lengths in order to get them in and out of the cache642 r = random.randint(0, 4)643 prompt = (" Hello " + str(r)) * (40 + r)644 n_prompt = (40 + r)*5 + 2645 n_predict = random.randint(1, 8)646 647 res = server.make_request(648 "POST",649 "/completion",650 data={651 "prompt": prompt,652 "n_predict": n_predict,653 },654 )655 656 assert res.status_code == 200657 assert "content" in res.body658 content = res.body["content"]659 assert isinstance(content, str)660 assert len(content) > 0661 662 assert type(res.body["has_new_line"]) == bool663 assert "timings" in res.body664 timings = res.body["timings"]665 666 assert "prompt_n" in timings and timings["prompt_n"] + timings["cache_n"] == n_prompt667 assert "predicted_n" in timings and timings["predicted_n"] == n_predict668 assert "tokens" in res.body and isinstance(res.body["tokens"], list)669 